Attention Mechanisms and Transformers

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 ii while looking at token ii 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.

The causal mask: what token 3 is allowed to see111111111111111k1k2k3k4k5q1q2q3q4q5Blanks are set to minus infinity before the softmax, so they receive exactly zero weight.
Mask after the scores and before the softmax — masking the weights afterwards leaves the rows no longer summing to 1.

Part 1: masking

The causal mask

Position ii may attend to positions j≤ij \le i. Build it once and cache it:

Python
import torchdef causal_mask(T, device=None):    """True where attention is ALLOWED. Lower-triangular including the diagonal."""    return torch.tril(torch.ones(T, T, dtype=torch.bool, device=device))print(causal_mask(4).int())# tensor([[1, 0, 0, 0],#         [1, 1, 0, 0],#         [1, 1, 1, 0],#         [1, 1, 1, 1]])

Read row ii as "what query ii 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:

Python
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 −∞-\infty 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][2.0,\ 1.0,\ 3.0,\ 5.0]:

Approachw1w_1w2w_2w3w_3w4w_4Row sum
No mask0.04140.01520.11240.83101.000
Softmax then zero the future0.04140.01520.112400.169
−∞-\infty then softmax0.24470.09000.665201.000

Check the arithmetic on the correct row: e2=7.389e^{2} = 7.389, e1=2.718e^{1} = 2.718, e3=20.086e^{3} = 20.086, e−∞=0e^{-\infty} = 0, sum =30.193= 30.193. Then 7.389/30.193=0.24477.389/30.193 = 0.2447, 2.718/30.193=0.09002.718/30.193 = 0.0900, 20.086/30.193=0.665220.086/30.193 = 0.6652. Sum =0.9999= 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\mathrm{softmax} of all-−∞-\infty is 0/00/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.

Python
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)(B, 1, 1, T) is deliberate. It broadcasts against scores of shape (B,h,Tq,Tk)(B, h, T_q, T_k): 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:

Text
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:

Python
def make_target_mask(tgt, pad_id=0):    B, T = tgt.shape    pad = (tgt != pad_id).unsqueeze(1).unsqueeze(2)          # (B, 1, 1, T)    causal = torch.tril(        torch.ones(T, T, dtype=torch.bool, device=tgt.device)    ).unsqueeze(0).unsqueeze(0)                              # (1, 1, T, T)    return pad & causal                                      # (B, 1, T, T)

For a target [5, 12, 9, 0] with pad_id = 0, the combined mask is:

Text
          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 it

Three masks are needed in a full encoder-decoder model, and mixing them up is a common source of quiet failure:

MaskUsed inShapeBlocks
Source paddingEncoder self-attention(B,1,1,Tsrc)(B,1,1,T_{src})Padded source columns
Target padding + causalDecoder self-attention(B,1,Ttgt,Ttgt)(B,1,T_{tgt},T_{tgt})Padded target columns and all future positions
Source paddingDecoder cross-attention(B,1,1,Tsrc)(B,1,1,T_{src})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 ii to source positions ≤i\le 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:

Python
@torch.no_grad()def test_no_future_leak(model, seq_len=8, vocab=100):    model.eval()    x = torch.randint(1, vocab, (1, seq_len))    out_a = model(x)    x2 = x.clone()    x2[0, -1] = (x2[0, -1] + 1) % vocab          # change ONLY the last token    out_b = model(x2)    # every position before the last must be bit-identical    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.

lr(t)=dmodel−0.5⋅min⁡ ⁣(t−0.5, t⋅warmup−1.5)lr(t) = d_{model}^{-0.5}\cdot\min\!\left(t^{-0.5},\ t \cdot \text{warmup}^{-1.5}\right)

Work out the actual values for dmodel=512d_{model}=512, warmup =4000=4000. First, 512−0.5=1/22.627=0.044194512^{-0.5} = 1/22.627 = 0.044194 and 40001.5=252,9824000^{1.5} = 252{,}982.

Step ttt−0.5t^{-0.5}t⋅4000−1.5t\cdot 4000^{-1.5}minLearning rate
1000.10000.0003950.0003951.75×10−51.75\times10^{-5}
1,0000.03160.003950.003951.75×10−41.75\times10^{-4}
4,0000.015810.015810.015816.99×10−4\mathbf{6.99\times10^{-4}} (peak)
16,0000.007910.06320.007913.49×10−43.49\times10^{-4}
100,0000.003160.3950.003161.40×10−41.40\times10^{-4}

The two branches cross exactly at t=4000t = 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.

Python
import mathfrom torch.optim.lr_scheduler import LambdaLRdef noam_schedule(optimizer, d_model, warmup=4000):    def fn(step):        step = max(step, 1)        return (d_model ** -0.5) * min(step ** -0.5, step * warmup ** -1.5)    return LambdaLR(optimizer, fn)optimizer = torch.optim.Adam(    model.parameters(), lr=1.0,          # lr=1.0 because the lambda IS the rate    betas=(0.9, 0.98), eps=1e-9,)scheduler = 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−310^{-3} and every rate above is a thousand times too small.

The betas are also non-default. Adam's usual β2=0.999\beta_2 = 0.999 averages the squared gradient over roughly 1000 steps; β2=0.98\beta_2 = 0.98 shortens that to about 50, which adapts faster to the rapidly changing gradient scale of early transformer training.

ScheduleShapeBest for
Noam (inverse sqrt)Linear warmup, then t−0.5t^{-0.5} decayTraining from scratch with no fixed step budget
Linear warmup + linear decayWarmup, then straight line to zeroFine-tuning; the default in Hugging Face
Warmup + cosine decayWarmup, then half a cosine to a small floorLarge pre-training runs with a known total step count
ConstantFlatNothing, 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.

Python
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

3.02+12.02=9+144=153=12.369\sqrt{3.0^2 + 12.0^2} = \sqrt{9 + 144} = \sqrt{153} = 12.369

With max_norm=1.0, every gradient is multiplied by 1.0/12.369=0.08081.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.

Python
accum_steps = 4          # micro-batch 8 -> effective batch 32optimizer.zero_grad(set_to_none=True)for i, batch in enumerate(loader):    loss = compute_loss(model, batch)    (loss / accum_steps).backward()          # divide BEFORE backward    if (i + 1) % accum_steps == 0:        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)        optimizer.step()        scheduler.step()        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.

float32float16bfloat16
Exponent bits858
Mantissa bits23107
Largest value3.4×10383.4\times10^{38}65,5043.4×10383.4\times10^{38}
Smallest normal1.2×10−381.2\times10^{-38}6.1×10−56.1\times10^{-5}1.2×10−381.2\times10^{-38}
Loss scaling neededNoYesNo

The problem is that smallest-normal row. Transformer gradients are routinely around 10−710^{-7} or smaller. In fp16 anything below 6.1×10−56.1\times10^{-5} becomes subnormal and below about 6×10−86\times10^{-8} becomes exactly zero — the gradient silently disappears.

Loss scaling fixes it. Multiply the loss by a large constant SS before backward(); by linearity every gradient is scaled by SS too. Then divide the gradients by SS before the optimiser step. With S=65536S = 65536, a gradient of 1×10−81\times10^{-8} becomes 6.55×10−46.55\times10^{-4} — comfortably inside fp16's normal range. GradScaler chooses SS automatically, halving it whenever it detects an overflow and doubling it after a stretch of clean steps.

Python
from torch.amp import autocast, GradScalerscaler = GradScaler('cuda')for batch in loader:    optimizer.zero_grad(set_to_none=True)    with autocast('cuda', dtype=torch.float16):        loss = compute_loss(model, batch)    scaler.scale(loss).backward()    scaler.unscale_(optimizer)                                 # undo scaling FIRST    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)    # ...then clip    scaler.step(optimizer)    scaler.update()    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/655361/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−ε1 - \varepsilon and the remaining ε\varepsilon is spread over the rest of the vocabulary. For ε=0.1\varepsilon = 0.1 and a 32,000-token vocabulary, each incorrect token receives

0.132000−1=3.125×10−6\frac{0.1}{32000 - 1} = 3.125\times10^{-6}

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 +∞+\infty, 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

Python
def train_epoch(model, loader, optimizer, scheduler, scaler,                accum_steps=1, clip=1.0, pad_id=0, device='cuda'):    model.train()    total_loss, total_tokens = 0.0, 0    optimizer.zero_grad(set_to_none=True)    for i, (src, tgt) in enumerate(loader):        src, tgt = src.to(device), tgt.to(device)        tgt_in, tgt_out = tgt[:, :-1], tgt[:, 1:]        src_mask = padding_mask(src, pad_id)        tgt_mask = make_target_mask(tgt_in, pad_id)        with autocast('cuda', dtype=torch.float16):            logits = model(src, tgt_in, src_mask, tgt_mask)            loss = F.cross_entropy(                logits.reshape(-1, logits.size(-1)),                tgt_out.reshape(-1),                ignore_index=pad_id,                label_smoothing=0.1,            )        scaler.scale(loss / accum_steps).backward()        if (i + 1) % accum_steps == 0:            scaler.unscale_(optimizer)            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), clip)            scaler.step(optimizer)            scaler.update()            scheduler.step()            optimizer.zero_grad(set_to_none=True)        n = (tgt_out != pad_id).sum().item()        total_loss += loss.item() * n        total_tokens += n    return total_loss / total_tokens

Weighting 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

SignalHealthyWhat a deviation means
Initial loss≈ln⁡V\approx \ln V — 10.4 for a 32k vocabularyMuch higher: bad initialisation or missing normalisation. Much lower: labels are leaking into the input
Pre-clip gradient normSettles into a band, e.g. 0.3–1.5Steadily rising: learning rate too high. Near zero: dead layers or a saturated softmax
Fraction of steps clippedUnder ~5% after warmupConsistently high: reduce the learning rate rather than lowering the clip threshold
Loss scale (fp16)Stabilises around 10410^{4}–10510^{5}Repeatedly halving: overflow somewhere; suspect the attention softmax or an unmasked −1e9-1e9 in fp16
Attention entropyFalls from ln⁡T\ln T towards 1–3 natsStuck at ln⁡T\ln T: heads are uniform and doing nothing. At 0: saturated, check dk\sqrt{d_k}
Train vs validation gapGrows slowlyValidation 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.

ObservationCauseFix
Validation loss suspiciously good; generation is degenerateCausal mask inverted or absent; or decoder input not shiftedRun the leakage test; print the mask; confirm tgt[:, :-1] vs tgt[:, 1:]
Loss diverges in the first few hundred stepsNo warmup, or post-norm without warmupAdd warmup over ~4000 steps; consider pre-norm
Model does not improve at all with fp16Clipping applied before unscale_Call scaler.unscale_(optimizer) first
Effective learning rate 4× too high; loss noisyLoss not divided by accum_stepsDivide before backward()
Warmup ends far too earlyscheduler.step() outside the accumulation guardMove it next to optimizer.step()
Large fraction of loss coming from paddingignore_index not setPass 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.