Transformers — Compute Attention By Hand
What Gets Recomputed for One More Token
In one line
Autoregressive generation produces tokens one at a time. Without a cache, each time you produce one more token, the k·v of the whole preceding context is rebuilt from scratch. The KV cache removes that rebuilding — it leaves the formulas as they are and changes only the number of computations.
Why this was needed
You have probably already computed attention by hand. You compute the scores, apply the softmax, and take a weighted average of the values. Three lines.
The problem arises when you call those three lines inside the generation loop. To produce an answer of length 100, you call the model 100 times. Each time you call it, the model takes "the whole context so far" as input. And it builds q·k·v for every position of that context.
But the k·v of the first token built in the second call is exactly the same value as the one built in the first call. It is the same in the third call. It is the same in the hundredth call. You build the same value a hundred times and throw away ninety-nine.
This waste is hard to see by reading the code. That is because there is nothing wrong in the attention function. To make it visible, you have to count.
Why count multiplications instead of time
From a measurement like "I turned the cache on and it got faster" you can learn nothing. How much faster depends on the machine, the load, the batch size and the memory bandwidth, and even on the same machine a second measurement gives a different number.
The number of multiplications is different. If you run the same code on the same input, it is always the same number. And that number directly shows how it grows with length. That is why this lab does not measure time. Instead you build one function that does the multiplication and count inside it.
def mul(a, b):
global _MULS
_MULS += 1
return a * b
It is the honest method. If you pass every place where multiplication happens through this function, you can count who multiplied how many times and where. You count the amount of computation instead of guessing it.
The shape you get by counting
Keep just one layer and one head, and call the model dimension d and the current context length n. The weights that make q·k·v are three matrices of d by d.
One token without a cache:
- Build q·k·v for all the preceding positions — 3 times n times d times d
- The last query looks at n keys — n times d
- Mix n values with those weights — n times d
One token with a cache:
- Build q·k·v for just one new token — 3 times d times d
- The query looks at n keys — n times d
- Mix n values with those weights — n times d
The difference is only the first line. On the side without a cache, the projection cost grows in proportion to the context length, and on the cache side that cost is fixed regardless of length. Attention itself (the last two lines) grows in proportion to length on both sides — the cache does not remove that. What the cache removes is recomputation, not attention.
In this lab you count those numbers by hand and build a table by length. You do not copy down a multiple someone else wrote down; you write the number you counted yourself.
The cache holds because of the causal mask
Why is it all right to reuse the k·v of earlier positions as they are? How can you be sure that the values of earlier positions do not change when a new token is attached?
It is because of the causal mask. In the decoder of Attention Is All You Need, each position is made to look only at what comes before it. The k and v of position 3 are made from the input of position 3 alone, and there is no reason for the values of position 3 to change just because token 4 is attached after it.
If there were no mask and every position looked at every other, the cache would not hold. That is because every time something is attached at the end, the earlier representations change. So the KV cache is a property of a decoder-only structure, not an optimization usable everywhere.
An important conclusion comes out here. The cache is not an approximation. The answer does not change. It just does not rebuild the same value. So if the output differs when you turn the cache on and off, that is not a property of the cache but a defect in the implementation.
The memory the cache consumes
When there is something you save, there is also something you pay. The cache consumes memory. The number of elements is written as a product.
원소 수 = 2 (K 와 V) × 층 수 × 헤드 수 × 문맥 길이 × 헤드 차원
바이트 = 원소 수 × 자료형 한 원소의 바이트 수
What to note here is that the context length is multiplied in. If the length doubles, the cache doubles too. This is one of the reasons long contexts are expensive, and this table is also usually what decides the number of requests a single machine can handle at once.
In this lab you set the layers, heads and head dimension to small values you choose and compute only with those values. You do not use GB figures of real models — because they were not measured here. The moment you copy down an unmeasured number, the text loses its basis. That changing the data type changes only the last term, and that the length is multiplied in, both of these are visible directly in the table you build.
What it looks like in the field
First, you resend a long prompt every time and cannot use the cache. When you continue a conversation, if you resend the whole earlier content, the server side has no basis for continuing to use the cache. This is why Hugging Face's Cache strategies documentation separately explains how to carry the cache object around yourself between calls.
Second, the upper limit on the number of concurrent requests comes from memory, not computation. The cache is allocated separately for each request and grows in proportion to length. So a few long conversations take more room than dozens of short requests.
Third, if you hold the cache in lower precision, the output changes slightly. The cache itself is not an approximation, but writing to the cache less accurately is one. In this lab you round before putting into the cache and directly check whether that difference exceeds the tolerance.
Fourth, the first token and the tokens after it differ in character. The first computation that reads the whole prompt has an empty cache and so has nothing to save. The cache pays off from the second token. Hugging Face's Text generation documentation also explains generation as divided into these two phases.
Fifth, if you mix batches while using the cache, positions get misaligned. The order of rows in the cache is the token order. If you bundle and unbundle requests and join the rows wrongly, no error occurs and the text just becomes strange.
What really matters in practice
- Count the number of operations instead of time. Time measures the environment and counts measure the algorithm. What is expensive and why can be seen only in the counts.
- Pin down with a test that the output is the same with the cache on and off. It is a one-line test, but if it breaks, the cache implementation is silently wrong. When comparing, use a tolerance, not an equality sign.
- Distinguish what is saved from what still grows. The projection becomes fixed, but attention keeps growing in proportion to length. If you talk about them mixed together, you misjudge the cost of long contexts.
- Write the memory calculation as five multiplications. Layers, heads, length, head dimension and data type. You can see at a glance what changes when which term changes.
- Do not copy down numbers you have not measured. "It gets N times faster" comes from that person's machine. If you count it yourself with your own settings, there is nothing to copy.
What you will do in the next lab
You grow /root/work/tf-kv/kv.py one step at a time. You use only the standard library — the system Python of this Pod has no numpy (it exists only inside /opt/onnx-lab/bin/python) and no torch or transformers either. So every number that appears here was counted by the code you built.
You start with a multiplication function that counts multiplications and a dot product, and build in turn a function that projects one token, attention in which one query looks at the piled-up K·V, a version that runs without a cache, and a version that runs with a cache. Then you run the two versions on the same input and check with a tolerance that the outputs are the same values, and count the multiplications by length and make a table.
The last two steps are the point of this lab. If you round what you put into the cache, the outputs of the two methods no longer fall within the tolerance — the fact that the cache is not an approximation but writing to it less accurately is one comes out in numbers. And with the layers, heads, head dimension and data type you chose, you compute the number of elements and bytes the cache consumes, and leave it as a record together with the multiplication-count table. The grader actually imports your module, pokes at the functions with inputs of different sizes every time, and computes the values and the number of multiplications separately to check against them.