Skip to content

Masks: causal, padding, sliding window, additive

ModuleL5.2 · build · Python · Pass 5 · 1 to 2 h, plus your graded tests (rung R5)
You buildpython/tinyllm/xfmr/masks.py: causal_mask, padding_mask, sliding_window_mask, combine, to_additive; and your own oracle tests in python/tests/l5-2-masks/
Contractcourse/contracts/py/tinyllm/xfmr/masks.pyi
Testscourse/tests/L5.2/test_masks.py (what they check: section 4), golden values from torch 2.14.1 in course/fixtures/L5.2/masks_torch.npz (course/oracle/L5.2/masks_torch.py); your tests are graded by mutation, threshold 0.80 with every pitfall fault required
Needsnothing to call. Reading: S-M05 (logic: AND, implication), M09.2 stable softmax (an all −∞-\infty row gives zeros)
Used byL5.3 and L5.5 build every attention mask here · later L6.1 (GPT’s causal mask), L7.7 (sliding windows), L8.2 (q_offset while decoding with a KV cache), L9.3 and L9.4 (the same rules as kernel flags)
MilestoneMS-L5
Optional depthVaswani et al., “Attention Is All You Need” (2017), section 3.2.3 (masking in the decoder); Beltagy, Peters, Cohan, “Longformer” (2020) and Jiang et al., “Mistral 7B” (2023), section 2 (sliding-window attention)
  • A mask is a boolean matrix over (query, key) with True = may attend; rules compose by AND (test_hand_example_masks, test_combine_is_and_with_broadcasting).
  • Causal: query ii at absolute position qoff+iq_{\text{off}} + i sees keys j≤qoff+ij \le q_{\text{off}} + i. One offset turns the training triangle into a decode chunk aligned to the bottom right (test_causal_matches_torch_biases, test_chunked_decode_masks_are_rows_of_the_full_mask).
  • Causality is checkable to the bit: changing future tokens must leave past outputs bitwise unchanged (test_causality_is_bitwise).
  • A sliding window is causal AND local: at most window keys, the query’s own included (test_sliding_window_matches_flex_attention, test_window_counts).
  • The additive form is 00 and −∞-\infty, never a big negative number: a fully blocked row must give zero weights, not a uniform average over forbidden keys (test_fully_masked_row_is_all_minus_inf).
Terminal window
ol start L5.2 # stubs masks.py; prints your test path and rung (R5)
ol tests L5.2 # the course tests
# write your oracle tests in python/tests/l5-2-masks/, then:
ol check L5.2 # course tests and the mutation grade of your tests
ol diff L5.2 # after passing: your code against the reference

Scaled dot-product attention (L5.1) lets every position read every other position, which is exactly wrong for the model you are about to build. A language model is trained to predict token t+1t + 1 from tokens up to tt; if position tt can read position t+1t + 1, training loss falls to nearly zero by copying, and generation, which has no future to copy, produces noise. In a batch, short sequences are padded, and a query that reads padding learns from garbage. Long contexts (Part 7) read only a recent window. All three are the same mechanism: a boolean matrix of allowed (query, key) pairs, turned into 00 or −∞-\infty added to the scores. Getting the off-by-ones right here is the difference between a model that trains and one that cheats, and the decode path of Part 8 reuses the same function with an offset.

SymbolMeaningType / shape
TqT_q, TkT_knumber of queries and keysint
ii, jjquery row and key columnint
qoffq_{\text{off}} (q_offset)absolute position of query 0; query ii is at qoff+iq_{\text{off}} + iint ≥0\ge 0
MMa boolean mask, True = may attendbool[Tq, Tk] (or broadcastable)
ℓb\ell_breal length of sequence bb in a padded batchint
ww (window)sliding-window sizeint ≥1\ge 1
AAthe additive mask: 0 where MM is True, −∞-\infty where Falsefloat[Tq, Tk]
SSattention scores QK⊤/dQK^\top/\sqrt{d}float[..., Tq, Tk]

Every mask in the system means the same thing: entry (i,j)(i, j) True means “query ii may read key jj”. This is the convention of torch.nn.functional.scaled_dot_product_attention. Beware: torch.nn.MultiheadAttention uses the opposite (True = blocked). Booleans are the interface, so the functions reject float or integer masks rather than guessing which convention a 0/1 array meant.

Position pp may read positions ≤p\le p, itself included (it needs its own token to predict the next one). With queries at absolute positions qoff+iq_{\text{off}} + i and keys at jj:

Mij=[ j≤qoff+i ].M_{ij} = [\,j \le q_{\text{off}} + i\,].

During training Tq=Tk=TT_q = T_k = T and qoff=0q_{\text{off}} = 0: the lower triangle with the diagonal. During generation with a KV cache (L8.2), the keys of all TkT_k tokens so far are cached, and a step computes only the last TqT_q queries, at positions Tk−Tq,…,Tk−1T_k - T_q, \ldots, T_k - 1: qoff=Tk−Tqq_{\text{off}} = T_k - T_q, and the triangle is aligned to the bottom right, so the last query sees every key. Torch ships both alignments (causal_upper_left, causal_lower_right); one parameter covers both here. The key test: the chunk’s mask is exactly the last TqT_q rows of the full mask, so cached and recomputed attention agree.

A batch stacks sequences of lengths ℓb\ell_b padded to TT. The key-padding mask is Pbt=[t<ℓb]P_{bt} = [t < \ell_b], shape [B, T]. To use it on scores of shape [B, H, Tq, Tk], index it as P[:, None, None, :]: one row per sequence, the same for every head and every query. Rules compose by AND, with numpy broadcasting: combine(causal_mask(T), P[:, None, None, :]) has shape [B, 1, T, T] and opens (i,j)(i, j) only if j≤ij \le i and j<ℓbj < \ell_b. Padded queries still produce outputs; the loss ignores them, so they need no mask.

Attention cost grows with TqTkT_q T_k. A sliding window keeps only the last ww keys:

Mij=[ j≤qoff+i ]∧[ (qoff+i)−j<w ].M_{ij} = [\,j \le q_{\text{off}} + i\,] \wedge [\,(q_{\text{off}} + i) - j < w\,].

Each query sees min⁡(w,p+1)\min(w, p + 1) keys. A stack of LL windowed layers still has a receptive field of about LwL w tokens, which is how Mistral reads long contexts with a 4096-token window. A window of at least TT is the causal mask; a window without the causal half lets queries read w−1w - 1 future tokens.

The softmax takes scores, not booleans. Add Aij=0A_{ij} = 0 where allowed and −∞-\infty where blocked: e−∞=0e^{-\infty} = 0 exactly, so blocked keys get weight exactly 0, and allowed scores are untouched (adding 0.00.0 changes nothing, not even the sign of a zero). If every key of a row is blocked, the row is all −∞-\infty; the stable softmax of M09.2 defines that as zeros. A “large negative” such as −109-10^9 looks equivalent but is not: in an all-blocked row every entry is −109-10^9, the softmax subtracts the maximum, and the row becomes uniform: the query averages over keys it was forbidden to read. In float16, −109-10^9 is not even representable (it rounds to −∞-\infty anyway). Because −∞-\infty contributes exact zeros, causality holds to the bit: perturbing a future token changes its score, but the score is masked to −∞-\infty and its value vector is multiplied by an exact 0.

Causal, T=3T = 3, and a decode chunk of 2 queries over 5 keys (qoff=3q_{\text{off}} = 3: queries at positions 3 and 4):

causal_mask(3)=(100110111),causal_mask(2,5,3)=(1111011111).\text{causal\_mask}(3) = \begin{pmatrix} 1&0&0\\1&1&0\\1&1&1 \end{pmatrix}, \qquad \text{causal\_mask}(2, 5, 3) = \begin{pmatrix} 1&1&1&1&0\\1&1&1&1&1 \end{pmatrix}.

The chunk is rows 3 and 4 of causal_mask(5). A window of 2 over 4 positions is a band: row pp opens p−1p - 1 and pp:

sliding_window_mask(4,4,2)=(1000110001100011).\text{sliding\_window\_mask}(4, 4, 2) = \begin{pmatrix} 1&0&0&0\\1&1&0&0\\0&1&1&0\\0&0&1&1 \end{pmatrix}.

Padding with lengths (3,1)(3, 1) and T=3T = 3 gives rows (1,1,1)(1, 1, 1) and (1,0,0)(1, 0, 0). Combined with the causal mask (padding broadcast over queries), sequence 0 keeps the triangle and sequence 1 opens only key 0 in every row.

Now the additive form on one row: scores (2,1,5)(2, 1, 5) with the third key blocked. A=(0,0,−∞)A = (0, 0, -\infty), the sum is (2,1,−∞)(2, 1, -\infty), and the softmax is (e2,e1,0)/(e2+e1)=(e/(e+1),1/(e+1),0)=(0.731059,0.268941,0)(e^2, e^1, 0)/(e^2 + e^1) = (e/(e + 1), 1/(e + 1), 0) = (0.731059, 0.268941, 0). The blocked key had the largest score and still gets exactly 0. These are the first two tests.

def causal_mask(Tq: int, Tk: int | None = None, q_offset: int = 0) -> NDArray: ... # bool [Tq, Tk]
def padding_mask(lengths: ArrayLike, T: int) -> NDArray: ... # bool [B, T]
def sliding_window_mask(Tq: int, Tk: int, window: int, q_offset: int = 0) -> NDArray: ...
def combine(*masks: ArrayLike) -> NDArray: ... # AND, broadcast
def to_additive(mask: ArrayLike, dtype=np.float32) -> NDArray: ... # 0 / -inf

Tk defaults to q_offset + Tq. Everything raises ValueError on sizes below 1, negative offsets, lengths outside [0,T][0, T], non-bool masks, shapes that do not broadcast, or a non-float dtype.

TestKINDChecksWhy it matters downstream
test_hand_example_masksunit, smokethe section 3 matricesyou and the tests agree on the convention
test_hand_example_additive_softmaxunit, smoke(0.731059,0.268941,0)(0.731059, 0.268941, 0)blocked means weight exactly 0
test_causal_matches_torch_biasesgoldentorch’s upper-left and lower-right causal biasesprefill and decode alignment
test_attention_with_masks_matches_torch_sdpagoldenattention through causal, decode-chunk, and padded masks equals torchthe masks do what attention needs
test_sliding_window_matches_flex_attentiongoldentorch flex_attention’s window of 3L7.7’s Mistral window
test_causality_is_bitwisepropertyfuture perturbations leave past outputs bitwise equalno leak of the answer in training
test_chunked_decode_masks_are_rows_of_the_full_maskpropertychunk masks are the last rows of the full masksL8.2’s cache vs recompute test
test_window_countspropertymin⁡(w,p+1)\min(w, p + 1) keys per row; a wide window is causalthe window rule exactly
test_combine_is_and_with_broadcastingpropertyAND, order-free, [B, 1, T, T] from [T, T] and [B, 1, 1, T]one mask per batch for every head
test_padding_mask_edgesboundarylengths 0 and TTempty slots in a batch
test_fully_masked_row_is_all_minus_infboundaryan all-blocked row is all −∞-\infty and attends to nothingpadded queries never read padding
test_additive_dtypes_and_valuesunitfloat16/32/64 with exact 0.0 and −∞-\inftymasks join scores in their dtype
test_default_tk_follows_q_offsetunitcausal_mask(2, q_offset=3) is [2, 5]cached keys are not dropped
test_rejects_bad_argumentsboundarybad sizes, offsets, lengths, float masks, int dtypesthe convention mix-up fails loudly

The oracle is the definition written as two nested loops over (i,j)(i, j): compare causal_mask and sliding_window_mask with it on several shapes and offsets (including a decode step), check padding on lengths 0 and TT, check combine against &, check the exact 00 and −∞-\infty of to_additive (including an all-blocked row), and check causality through a few lines of numpy attention. Import only tinyllm.xfmr.masks.

PitfallSymptomCaught by
1. a strict j<ij < i causal maskposition 0 sees nothing; each token cannot see itselftest_hand_example_masks (mutant s01)
2. ignoring q_offset (top-left alignment for a decode chunk), or a default Tk = Tqcached decoding differs from recomputationtest_chunked_decode_masks_are_rows_of_the_full_mask (mutant s02), test_default_tk_follows_q_offset (mutant s09)
3. a window of w+1w + 1 keys (≤\le for <<)disagrees with Mistral checkpointstest_window_counts (mutant s03)
4. a window without the causal halfthe model reads w−1w - 1 future tokenstest_causality_is_bitwise (mutant s04)
5. padding open at t=ℓbt = \ell_bevery sequence reads one pad tokentest_padding_mask_edges (mutant s05)
6. a finite “minus infinity”an all-blocked row averages forbidden keystest_fully_masked_row_is_all_minus_inf (mutant s06)
7. the inverted convention (−∞-\infty where True)attention reads only what it should nottest_hand_example_additive_softmax (mutant s07)
8. combining by ORpadding re-opened by the causal ruletest_combine_is_and_with_broadcasting (mutant s08)
DirectionModuleHow it uses this
BackS-M05AND of conditions, and implication as a matrix (reading)
BackM09.2the softmax that maps an all −∞-\infty row to zeros (reading)
ForwardL5.1scaled_dot_product_attention(q, k, v, mask) takes these masks (True = attend)
ForwardL5.3multi-head attention broadcasts one mask over every head
ForwardL5.5the decoder’s mask is combine(causal_mask(T), padding)
ForwardL6.1GPT trains under causal_mask(T)
ForwardL7.7sliding-window attention in the modern block
ForwardL8.2the KV-cache step is causal_mask(Tq, Tk, q_offset=Tk - Tq)
ForwardL9.3the C attention kernel takes causal, q_offset, and window flags with these meanings
Your pieceProduction equivalentWhat it addsWhere to look
causal_mask with q_offsettorch.nn.attention.bias.causal_lower_righta lazy bias the fused kernels recognize, never materializedtorch/nn/attention/bias.py
sliding_window_maskFlexAttention mask_mod + BlockMaskarbitrary rules compiled into block-sparse kernels that skip fully blocked tilestorch/nn/attention/flex_attention.py
to_additiveFlashAttention’s causal and window_size argumentsthe mask is never a tensor; tiles above the diagonal are skippedflash_attn/flash_attn_interface.py
combine with paddingHF transformers masking_utilsone function builds causal, sliding, chunked, and padding masks per layer typesrc/transformers/masking_utils.py