Transformer Attention Math: Q/K/V, Softmax Weights, Masks, and KV Cache
Transformer Attention Math: Q/K/V, Softmax Weights, Masks, and KV Cache
Search
Ask the AI

Transformer Attention Math: Q/K/V, Softmax Weights, Masks, and KV Cache

Attention bugs often survive a shape check. A program can return a correctly sized tensor while reading future tokens, normalizing the wrong axis, or ignoring most of a cached prefix. This article follows one fixed three-token example from scores to outputs, then tests the same calculation during incremental decoding.

Experiment revision: 8 September 2026. Use the Attention and KV Cache Lab for this page. The older Deep Learning Math Lab linked elsewhere on the page uses different, unmasked attention inputs. Its heatmap cannot verify the causal example below. No trained model or GPU is required.

1. Read the Formula with Explicit Shapes

scores = (Q @ transpose_last_two_axes(K)) / sqrt(d_k)
weights = softmax(masked_scores, axis=keys)
context = weights @ V

For one head, Q has shape (Lq, dk), K is (Lk, dk), and V is (Lk, dv). Scores are (Lq, Lk); context is (Lq, dv). The square case Lq = Lk is only one case. Cached decoding below has one query and three keys.

The scaling argument assumes independent, zero-mean, unit-variance query/key components: the dot-product variance is dk, and division by its square root makes that variance one. This is a motivation, not a guarantee for learned activations. dk is the per-head dimension, not automatically the model width. See the definition in Attention Is All You Need.

Separate learned Q/K/V projections let the model form a query, a matching key, and a value to retrieve without requiring identical coordinates. Our fixture starts after those projections: its numbers are manually specified, not embeddings that demonstrate the meaning of “AI needs math.”

2. Multi-Head Data Flow

graph TD
    Input[Input: B, L, d_model] --> WQ[W_q projection and split]
    Input --> WK[W_k projection and split]
    Input --> WV[W_v projection and split]
    WQ --> Q[Q: B, h, L, d_k]
    WK --> K[K: B, h, L, d_k]
    WV --> V[V: B, h, L, d_v]
    Q --> Dot[Q times transposed K: B, h, L, L]
    K --> Dot
    Dot --> Scale[Divide by sqrt d_k]
    Scale --> Mask[Mask future keys]
    Mask --> Softmax[Softmax over keys]
    Softmax --> Read[Weights times V]
    V --> Read
    Read --> Context[Context: B, h, L, d_v]
    Context --> Concat[Concatenate: B, L, h times d_v]
    Concat --> Out[W_o projection: B, L, d_model]

Transpose only K’s last two axes in a batched implementation. With NumPy, K.T reverses every axis of a four-dimensional tensor; use K.swapaxes(-1, -2). Concatenation has width h * dv; the output projection maps it to dmodel. These widths are equal only when the chosen configuration makes them equal.

3. One Fixture, All Intermediate Matrices

Positions are numbered 0, 1, 2. Each query can read itself and earlier positions. Q/K have dimension 4 and V has dimension 2. In next-token training, the representation at input position i is used to predict the target at i+1; “current input” and “next target” must not be confused.

Complete runnable NumPy example
import numpy as np

# Fixed illustrative tensors, not embeddings from a trained model.
Q = np.array([[1.0, 0.5, -0.2, 0.1],
              [-0.5, 1.2, 0.8, -0.4],
              [0.2, -0.1, 1.5, 0.9]], dtype=np.float64)
K = np.array([[0.8, 0.4, -0.3, 0.0],
              [-0.2, 1.0, 0.5, -0.1],
              [0.1, -0.2, 1.1, 0.7]], dtype=np.float64)
V = np.array([[1.0, 0.0], [0.0, 1.0], [-1.0, -1.0]])

scores = (Q @ K.T) / np.sqrt(Q.shape[-1])
allow = np.arange(3)[None, :] <= np.arange(3)[:, None]
masked = np.where(allow, scores, -np.inf)
shifted = masked - masked.max(axis=-1, keepdims=True)
weights = np.exp(shifted)
weights /= weights.sum(axis=-1, keepdims=True)
context = weights @ V

if __name__ == "__main__":
    np.set_printoptions(precision=6, suppress=True)
    print("Scaled scores:\n", scores)
    print("Masked scores:\n", masked)
    print("Attention weights:\n", weights)
    print("Context:\n", context)

The first score is independently checkable: (1*0.8 + 0.5*0.4 + (-0.2)*(-0.3) + 0.1*0) / 2 = 0.53. The other dot products produce:

Scaled scores                 After causal masking
[[ 0.530  0.095 -0.075]        [[ 0.530   -inf   -inf]
 [-0.080  0.870  0.155]         [-0.080  0.870   -inf]
 [-0.165  0.260  1.160]]        [-0.165  0.260  1.160]]

Attention weights             Context = weights @ V
[[1.000000 0.000000 0.000000]  [[ 1.000000  0.000000]
 [0.278885 0.721115 0.000000]   [ 0.278885  0.721115]
 [0.158938 0.243109 0.597953]]  [-0.439015 -0.354843]]

For query 1, subtracting the largest visible score yields [-0.95, 0, -inf]. Its denominator is exp(-0.95) + 1, giving weights approximately 0.278885 and 0.721115. Query 2 also reads itself: its context is [w20 - w22, w21 - w22] because V2 is [-1,-1]. This explains both negative output components without invoking a trained language model.

Causal attention for the same three-token dk4 example; all three future-key cells are masked to zero
The chart is generated from the exact arrays above by plot_attention.py. Grey cells are forbidden reads, not missing data. These are attention weights, not gradients or evidence that a model understands language.

4. The One-Query Cache Trap

A full 3-by-3 lower triangle is correct for a query starting at position zero. Reusing a freshly created 1-by-3 lower triangle for the last cached query is not: that query is at absolute position 2, not position 0.

Mask Output (x, y)
Offset 2: keys 0 to 2 (-0.439, -0.355)
Wrong offset 0: key 0 (1.000, 0.000)

The wrong result has a maximum component error of 1.439015. Its row still sums to one and its shape is still correct. A row-sum assertion alone therefore misses this bug. Build visibility from absolute positions: key_position <= query_start + local_query_position. In this full-prefix lab, query_start is the number of already cached tokens.

This also matters when translating code to a library. The PyTorch 2.14 SDPA documentation specifies upper-left alignment for its non-square causal mask; its boolean attention mask uses True for visible entries. Do not assume another API’s padding-mask convention is the same. This article compares documented semantics; it does not claim to have run a PyTorch parity test.

The checked implementation rejects an all-masked query row. Otherwise subtracting the row maximum can become -inf - -inf and produce NaN. This rejection is our explicit lab contract, not a universal claim about framework behavior. Combine causal visibility and padding visibility logically, and decide what an entirely padded query should do before applying softmax.

5. What the Executed Cache Tests Show

The reference run used Python 3.13.9, NumPy 2.3.5 and float64. The package contains source hashes, full-precision CSVs and the audit JSON. A separate 50-digit Decimal implementation uses loops and Decimal.exp, starting from the exact binary64 fixture values.

Check Result
Decimal output 5.6e-17
Token cache [1,1,1] 5.6e-17
Chunk cache [2,1] 0
Batched cache [2,3,2] 2.2e-16
Future K/V change 0
Invalid inputs 12 rejected

The first five results are maximum absolute output errors. Decimal checks cover all 9 scores, 9 weights and 6 outputs; the maximum score error is 1.1e-16. The batched test uses B=2, H=3, L=7, dk=4, dv=5 and seed 42. Both caches append K and V with the correct prefix offset. Changing future K/V leaves the first two outputs unchanged. Rejected inputs include nonfinite values, all-masked rows, incorrect shapes and invalid offsets.

Additional negative controls deliberately reverse the causal mask or normalize columns. They produce forbidden future-weight mass of 0.872569 and a maximum row-sum error of 0.596413, respectively. The tests detect both. A failed cache append also leaves the previous K/V unchanged.

The cache stores actual appended arrays and checks that their prefix is unchanged. It does not simulate a complete Transformer: Q/K/V are already supplied. Repeated NumPy concatenation copies history and is not an efficient serving design. No speedup, GPU kernel agreement, trained-model accuracy or language-generation quality was measured.

6. Estimate KV Memory Before Calling It an OOM Problem

For equal K/V head dimensions and dtype, count their raw stored elements:

bytes = 2 * layers * batch * tokens * kv_heads * head_dim * bytes_per_element

For 32 layers, batch 1, 10,000 tokens, head dimension 128 and 2-byte elements, the same formula gives the following estimates, not measured GPU allocations:

KV heads Raw GiB Interpretation
32 4.8828 MHA with 32 query heads
8 1.2207 GQA with 32 query heads
1 0.1526 MQA with 32 query heads

The script checks the formula against a small pair of real NumPy fp16 arrays: 480 bytes. It does not allocate the multi-GiB examples. For comparison, exactly seven billion 2-byte parameters occupy about 13.04 GiB before overhead. Thus this 10K-token MHA example does not exceed those weight bytes. A conclusion about OOM also needs activations, allocator behavior, metadata, batch size and the actual architecture.

PagedAttention addresses cache allocation waste and sharing; paging alone does not reduce the raw element count above. GQA and MQA reduce the number of KV heads. These are distinct mechanisms, not interchangeable claims that every cache needs one huge contiguous block. Neither is implemented by this small lab.

7. Reproduce and Inspect a Failure

Extract the lab ZIP, then run from its directory:

python3 -m venv .venv
source .venv/bin/activate
python -m pip install -r requirements.txt
python attention_example.py
python audit_attention.py --out results

Success ends with ATTENTION_AUDIT_OK. The default results directory does not overwrite the packaged reference. To regenerate the chart, install requirements-plot.txt and run plot_attention.py. CSVs retain 17 significant digits; expect tolerance-based agreement rather than identical final floating-point bits across machines.

A useful exercise is to change the cache query offset from 2 to 0 and reproduce the wrong [1,0] output above. Then restore it and change only the final V: the earlier causal outputs must remain unchanged. These controlled edits test visibility, rather than merely producing a plausible-looking heatmap.

8. Scope and Next Step

Cache reuse depends on an unchanged prefix, model parameters and position conventions during causal inference. Editing a prefix or changing its positions can invalidate cached values; dropout and training require separate treatment. Hugging Face’s cache explanation describes how per-layer cached keys and values are combined with the new token’s entries.

This article establishes a reproducible arithmetic and masking example, not causal interpretability or production readiness. To continue, connect these forward values to the backpropagation computation-graph experiment, where finite differences test derivatives rather than attention visibility.

Leave a Reply

Scroll down