Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

What purpose does dropout serve when applied within the attention mechanism?


What you need to know

Where it sits

Text
weights = softmax(QKᵀ / sqrt(d) + mask)weights = dropout(weights, p)        # training onlyout     = weights · V
Python
import torchtorch.manual_seed(0)attn = torch.full((3, 4), 0.25)             # each row sums to 1drop = torch.nn.Dropout(p=0.25)drop.train()print(drop(attn))                           # some zeros, survivors become 0.3333print(drop(attn).sum(-1))                   # e.g. tensor([1.3333, 1.0000, 1.3333])drop.eval()print(torch.equal(drop(attn), attn))        # True: off at inference

Rows sum to 1 only on average. Each training step, a head sees a slightly different mix, which acts like training many attention patterns at once.

What it regularises

A head that routes everything through one position is fragile: if that token is noisy or absent at test time, the head fails. Dropping that link now and then forces the head to also use other evidence. It is the same "no co-adaptation" argument as ordinary dropout, applied to routing decisions instead of features.

Current practice

SettingAttention dropout
Original Transformer, BERT, GPT-20.1
Large LLM pretraining (Llama family and most since)0.0
Fine-tuning a small model on a small datasetoften 0.1
LoRA fine-tuningbase model dropout 0; LoRA's own dropout, often 0.05–0.1

Why zero at scale: pretraining sees trillions of mostly unique tokens once, so overfitting is not the main risk, and dropout slows training and wastes capacity.

When you call PyTorch's F.scaled_dot_product_attention directly, it has a dropout_p argument and does not know whether the model is in training mode. You must pass dropout_p=p if self.training else 0.0 yourself.

A real-life example

An e-commerce company fine-tunes a BERT-style cross-encoder for search ranking on 50,000 labelled query-product pairs. With so little data, the model soon memorises the training pairs: on the validation set, ranking quality peaks after one epoch and then falls. They keep the default attention dropout of 0.1 and add early stopping, and the validation curve flattens instead of falling.

Later, the ranking scores in production look noisy — the same query and product get slightly different scores on each request. The cause: their serving code called F.scaled_dot_product_attention with dropout_p=0.1 hard-coded. model.eval() does not affect a number passed as an argument. Setting it to 0 at inference made scores deterministic again.

Follow-up questions to expect

  • "Why not renormalise rows after dropout?" — It would change the expected output; the 1/(1-p) scaling already preserves the expectation, which is what matters for training.
  • "Does FlashAttention support dropout?" — Yes; the kernel generates the dropout mask inside the tiles and regenerates it in the backward pass, so no mask is stored.
  • "Is attention dropout the same as dropping tokens?" — No. It drops individual query-to-key links, differently for each head and query, not whole tokens.