An attention score combines a query with every permitted key; a mask must encode exactly which key positions may influence each query position.
Scaled dot-product attention and mask contracts
Track four axes
For a batched multi-head implementation, queries have batch, head, query-position and feature axes; keys and values have batch, head, key-position and feature axes. Scores occupy batch by head by query-position by key-position. Scale dot products by the square root of the query/key feature width before softmax so score magnitude does not grow unchecked with width. The resulting weights mix value vectors, not key vectors. Tensor contracts should assert every axis before training.
Specify mask meaning per API
In one common scaled-dot-product API, a Boolean True means the key may participate; another multi-head API uses True to mean padding should be ignored. Carry a named allow-mask at the boundary and convert it deliberately when switching APIs. A padding mask blocks invalid key positions; a causal mask blocks keys to the right of the query. Combine them by logical AND. A key mask alone does not guarantee padded query outputs are zero, so zero or exclude those outputs before downstream pooling.
Avoid rows with no valid keys
Softmax over all negative-infinity scores is undefined in a simple reference implementation and may produce NaN. A request with zero valid tokens should fail validation or use an explicit safe fallback. With causal attention and right padding, valid query positions normally retain at least their own key. Left padding, cached decoding and cross-attention need separate tests, because their query and key lengths or alignments can differ.
Control dropout at evaluation
Some low-level attention functions apply the passed dropout probability even when the containing model is in evaluation mode. Pass zero explicitly during evaluation or serving. Compare deterministic outputs for the same input twice, and confirm no output at position p changes when only a future token changes. Train/eval mode checks apply to custom attention calls as well as ordinary dropout modules.
Budget sequence length
Dense score computation scales quadratically with sequence length and linearly with heads and batch size; specialized kernels may avoid storing the entire score matrix, but arithmetic still grows with attended pairs. Profile the longest accepted sequence rather than the median. A mask does not automatically make a dense kernel sparse. Restrict length, choose a supported efficient kernel or change the attention pattern when the cost exceeds the product budget.
Implementation
import torch
import torch.nn.functional as functional
torch.manual_seed(61)
receipt_tokens = torch.randn(1, 1, 4, 8)
valid_tokens = torch.tensor([[True, True, True, False]])
sequence_length = receipt_tokens.size(-2)
causal_allow = torch.ones(sequence_length, sequence_length,
dtype=torch.bool).tril()
key_allow = valid_tokens[:, None, None, :]
allow_mask = causal_allow[None, None, :, :] & key_allow
attended = functional.scaled_dot_product_attention(
receipt_tokens, receipt_tokens, receipt_tokens,
attn_mask=allow_mask, dropout_p=0.0)
attended = attended * valid_tokens[:, None, :, None]
assert attended.shape == receipt_tokens.shape
assert torch.count_nonzero(attended[:, :, -1, :]) == 0
assert not allow_mask[0, 0, 0, 1]
assert not allow_mask[0, 0, 2, 3]Performance and operating cost
A dense attention layer computes roughly O(B × H × L × S × D) work for batch B, heads H, query length L, key length S and head width D. A naive implementation also stores O(B × H × L × S) scores. Kernel choice can reduce stored intermediates, but padding still wastes work unless batching or the kernel exploits lengths. Validate masks before expensive model execution.
Common Mistakes
- Do not carry a Boolean mask between APIs without checking which truth value means allowed.
- Do not assume padding keys also zero padded query outputs.
- Do not rely on model.eval() alone to disable a dropout probability passed to a low-level attention function.
Read next
- Tensor contracts: shape, dtype, device and mask
- Training and validation modes: measure the model you will serve
- Causal target shifts and padding-aware loss
- Unicode and tokenization: preserve meaning at the text boundary
- Project: audit a next-event attention decoder
Continue the workflow: Autoregressive KV cache position and mask parity.
Continue the workflow: Vision transformer patch grids and token contracts.
Continue the workflow: Sparse expert routers and per-batch capacity accounting.
Continue the workflow: Encoder-decoder cross-attention and source masks.
