Transformers — Compute Attention By Hand
Count the Multiplications a KV Cache Saves
Goal
Confirm what the KV cache saves in autoregressive generation by counting the multiplications directly. You build a multiplication function that counts multiplications, build a version that runs without a cache and a version that runs with a cache, confirm with a tolerance that the outputs of the two versions are the same values, and then count the multiplications by context length and make a table. At the end you measure that this sameness breaks when the cache is held in lower precision, and with the settings you chose, compute the number of elements and bytes the cache consumes and leave it as a record.
Why it matters
Computing one attention by hand and calling that attention inside the generation loop are different stories. To produce an answer of length 100, you call the model 100 times, and without a cache, each time you call it, the k·v of the whole preceding context is rebuilt from scratch. That means you build the k·v of the first token a hundred times and throw away ninety-nine.
This waste is not visible by reading the code. That is because there is nothing wrong in the attention function itself. To make it visible, you have to count. And you must not measure time — time varies with the machine and the load, but the number of multiplications is always the same for the same input and directly shows how it grows with length.
This lab does not call a real model. The system Python of this Pod has no numpy (it exists only inside /opt/onnx-lab/bin/python) and no torch or transformers either. You build the same structure with the standard library alone and use only the numbers measured here. So statements like "a real model consumes N GB" or "it gets N times faster" are not made here.
The grader does not trust the explanations you wrote down. It 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. The sizes change on every run, so you cannot memorize values and plug them in.
Steps
- In /root/work/tf-kv/kv.py, create
CONFIGandreset_muls(),muls(),mul(a, b)anddot(u, v).mulmultiplies and raises the count by one, anddotmultiplies only through thatmul. - Add
project(x, W)so that it passes one token through one set of weights.Wis a matrix whose rows each have lengthlen(x), and the length of the list returned islen(W). - Add
attend(q, K, V)so that one query looks at the whole piled-up K·V. Divide the scores by√len(q), apply the softmax, and take the weighted average of V. - Add
step_nocache(xs, Wq, Wk, Wv)so that it produces one output without a cache. Rebuild q·k·v for all the preceding tokens and attend with the last query. - Add
new_cache(),store(cache, k, v, digits=None)andstep_cached(x, Wq, Wk, Wv, cache, digits=None)so that they produce the same output with a cache. Compute only one new token and append one row to K·V. - Add
ATOL,RTOL,close(got, want)andcompare(xs, Wq, Wk, Wv, digits=None)so that they compare whether the outputs of the two methods are the same values. Do not compare with an equality sign. - Add
mul_table(xs, Wq, Wk, Wv)so that it counts the multiplications of the two methods for context lengths from 1 tolen(xs). The return value is a list of(문맥길이, 캐시없음, 캐시있음)tuples (the placeholders are the context length, without a cache and with a cache). - Make a table with 12 tokens, compute the memory the cache consumes, and record it in /root/work/tf-kv/kv_report.json and /root/work/tf-kv/kv_report.md.
Notes
- Execution contract: the grader imports
/root/work/tf-kv/kv.pyas a Python module and usesCONFIG,reset_muls,muls,mul,dot,project,attend,step_nocache,new_cache,store,step_cached,ATOL,RTOL,close,compareandmul_tabledirectly. It does not run it as a script, soif __name__ == "__main__"is not needed. CONFIGis a dictionary with the four keys{"layers": ..., "heads": ..., "head_dim": ..., "dtype_bytes": ...}. You choose the values. The ranges are 2 to 8 layers, 2 to 8 heads, an even head dimension from 4 to 32, and data type bytes of 1, 2 or 4. You do not need to imitate the numbers of a real model — the memory calculation in step 8 uses only these values you chose.mul(a, b)returnsa * band raises the counter inside the module by 1.reset_muls()sets the counter to 0 andmuls()returns the current value. The count comes out right only if you pass every place where multiplication happens throughmul.dot(u, v)is the dot product. The multiplications happenlen(u)times. Additions, divisions and exponentials are not counted.- The multiplications of
project(x, W)numberlen(W)timeslen(x). A row ofWmakes one cell of the output — if you swap rows and cells, both the values and the counts go off. - The multiplications of
attend(q, K, V)arelen(K)timeslen(q)on the score side andlen(V)timeslen(V[0])on the weighted-sum side. Division and exponential are not multiplications, so they are not counted. For the softmax, take the exponential after subtracting the maximum. step_nocache(xs, Wq, Wk, Wv)builds q·k·v for all the tokens ofxsand looks at the whole K·V with the last query. Building everything even though only one query is used is what the no-cache state looks like.new_cache()returns{"K": [], "V": []}.store(cache, k, v, digits=None)appends one row to each of K and V — it does not overwrite. Ifdigitsis given, each component is rounded to that number of decimal places before being put in.step_cached(x, Wq, Wk, Wv, cache, digits=None)projects only the one new token, puts it into the cache withstore, and then looks at the whole cache with that query. If you attend before putting it in, the new token cannot see itself.- Set
ATOL = 1e-9andRTOL = 1e-6, andclose(got, want)isabs(got - want) <= ATOL + RTOL * abs(want). Write in a comment why this width. A Python float is IEEE 754 double precision, so it has about 15 significant digits, and even if you just change the order of adding in a dot product and softmax of this scale, the relative error stays around 1e-12. compare(xs, Wq, Wk, Wv, digits=None)runs the two methods while growing the length one step at a time from 1 tolen(xs)and returns{"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]}(the placeholders are a real number and a true/false value).max_gapis the largest of the absolute differences between components, andall_closeis whether every component passedclose.mul_table(xs, Wq, Wk, Wv)resets the counter just before measuring at each length. The cache side must keep using one cache continuously — if you create a new one for every row, that is not a cache.- Step 8 sets
d = CONFIG["head_dim"]and measures with 12 tokens and three weight matrices ofdbyd. It uses 2 forround_digits. - The token vectors and weights in step 8 are made with
((i * 7 + j * 5 + salt * 3) % 13 - 6) / 7.0.iis the row number andjis the cell number, andsaltis 1 for the token vector, 2 for Wq, 4 for Wk and 6 for Wv. The multiplication counts and memory come only from the sizes, not the values, so the table is the same whatever values you use, but the rounding experiment depends on the values — dividing by 7 is so that the values do not come out exact at the second decimal place. If you use only values that come out exact, rounding to the second decimal place leaves the values as they are, so the effect of rounding is not visible. - The memory is
원소 수 = 2 × layers × heads × 길이 × head_dimand바이트 = 원소 수 × dtype_bytes(the Korean words mean number of elements, length and bytes). The length is 12. - This Pod has no internet.
pip installdoes not work and there is no torch or transformers. numpy exists only inside/opt/onnx-lab/bin/python, soimport numpydoes not work in the system Python.mathalone is enough. - Do not measure time. A number measured with
timevaries with the machine and the load and cannot be used for judging. What this lab measures is the number of multiplications. - Official documents: Attention Is All You Need · Hugging Face — Cache strategies · Hugging Face — Text generation · Python — math
- Common mistakes: multiplying directly without going through
mulso that the count stays 0, swapping rows and cells inproject, not dividing the scores by√d, overwriting the cache, attending before putting into the cache, not resetting the counter inmul_table, and creating a new cache for every row.
A multiplication that counts multiplications
In /root/work/tf-kv/kv.py, create CONFIG and reset_muls(), muls(), mul(a, b) and dot(u, v). CONFIG is a dictionary with the four keys layers, heads, head_dim and dtype_bytes, and you choose the values (2 to 8 layers, 2 to 8 heads, an even head dimension from 4 to 32, and data type bytes of 1, 2 or 4). mul multiplies and raises the counter by 1, and dot multiplies only through that mul.
The counter can be a single integer in the module. To change it inside a function, you need global. If dot multiplies directly, as in sum(a * b for a, b in zip(u, v)), it counts nothing, so be sure to go through mul. Additions are not counted — the value of matrix computation comes from the multiplication side. Do not put in code that measures time.
Project one token
Add project(x, W). W is a matrix whose rows each have length len(x), and one row makes one cell of the output. The length of the list returned is len(W), and the multiplications happen len(W) times len(x).
You are done by calling the dot you built earlier once per row. It can be written in one line. If you swap rows and cells, only the values are wrong for a square matrix and even the length goes off for a rectangular one, so first check that the length of the returned list is len(W). q, k and v are all made with this one function.
One query looks at the whole cache
Add attend(q, K, V). Divide the dot product of q and each k by √len(q) to get the scores, subtract the maximum, take the exponential to apply the softmax, and take the weighted average of V with those weights. The multiplications are len(K) times len(q) on the score side and len(V) times len(V[0]) on the weighted-sum side.
Division and exponential are not multiplications, so do not wrap them in mul — if you do, the count goes off. Only the multiplications of the weighted sum go through mul. The length of one row of V may differ from the length of q, so set the output list by len(V[0]). If K has only one row, the weight is the single value 1, so the output must equal V[0].
Without a cache — recompute everything before
Add step_nocache(xs, Wq, Wk, Wv). Build q·k·v for all the tokens of xs and look at the whole K·V with the last query. Building everything even though only one query is used is what the no-cache state looks like.
Three list comprehensions and one attend are enough. Just because only the last query is used, you must not project only the last token — the keys and values must exist for all the earlier tokens, and the point of this step is the fact that you rebuild them every time. The number of multiplications grows in proportion to the context length.
With a cache — append only one row
Add new_cache(), store(cache, k, v, digits=None) and step_cached(x, Wq, Wk, Wv, cache, digits=None). new_cache() returns {"K": [], "V": []}, store appends one row to each of K and V (rounding to that number of decimal places if digits is given), and step_cached projects only the one new token, puts it into the cache, and then looks at the whole cache with that query.
Use append, not overwriting. The earlier rows do not change in value when a new token is attached — because of the causal mask, each position looks only at what comes before it, and that is why the cache holds. If you attend before putting it in, the new token cannot see itself and the answer differs from the no-cache version. In the number of multiplications, the projection side is fixed regardless of context length.
Are the outputs of the two methods the same value
Add ATOL = 1e-9, RTOL = 1e-6, close(got, want) and compare(xs, Wq, Wk, Wv, digits=None). close is abs(got - want) <= ATOL + RTOL * abs(want), and you write in a comment why this width. compare runs the two methods while growing the length one step at a time from 1 and returns {"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]} (the placeholders are a real number and a true/false value).
Do not compare with an equality sign. If the two methods happen to add in the same order, they may come out equal down to the bit, but if the order of adding differs even slightly, the last digit wobbles — that is why the right test is a tolerance. Pass digits on to step_cached as it is. If you round before putting in, all_close becomes false, and that is what this step is meant to show. max_gap is the largest difference across all lengths and all components.
Count the multiplications by length
Add mul_table(xs, Wq, Wk, Wv). Count the multiplications of the two methods for context lengths from 1 to len(xs) and return a list of (문맥길이, 캐시없음, 캐시있음) tuples (the placeholders are the context length, without a cache and with a cache). Reset the counter just before each measurement, and the cache side keeps using one cache continuously.
Call reset_muls() twice — once before measuring the no-cache version and once before measuring the cache version. If you do not reset, the earlier numbers get mixed into the later ones. If you create a new cache for every row, the cache-side numbers grow like the no-cache side — that is not a cache. In the table, the no-cache side grows in proportion to length, and on the cache side only the attention share grows.
Write down what is saved and what is paid together
Set d = CONFIG["head_dim"] and run mul_table and compare with 12 tokens and three weight matrices of d by d. Make the values with ((i * 7 + j * 5 + salt * 3) % 13 - 6) / 7.0, where salt is 1 for the token vector, 2 for Wq, 4 for Wk and 6 for Wv. Run compare once without rounding and once with digits=2. Then write layers, heads, head_dim, dtype_bytes, tokens, table, total_nocache, total_cached, saved_muls, cache_elems, cache_bytes, round_digits, max_gap_exact, all_close_exact, max_gap_rounded and all_close_rounded in /root/work/tf-kv/kv_report.json, and write /root/work/tf-kv/kv_report.md in the four sections ## 무엇을 쟀나 ## 캐시 없이 하면 무엇을 다시 계산하나 ## 캐시가 먹는 메모리 ## 두 방식의 출력이 같은가 (the Korean headings mean "What was measured", "What is recomputed without a cache", "The memory the cache consumes" and "Are the outputs of the two methods the same").
Do not write the numbers by hand; fill them in with values obtained by actually running your own code. total_nocache and total_cached are the sums of the cells of the table, and saved_muls is their difference. cache_elems is 2 × layers × heads × 12 × head_dim, and cache_bytes is that multiplied by dtype_bytes. all_close_exact must be true and all_close_rounded must be false — the cache itself is not an approximation, but writing to the cache less accurately is one. Dividing by 7 in the formula that makes the values is so that the values do not come out exact at the second decimal place. If you use only values that come out exact, rounding leaves the values as they are and the experiment itself does not hold. Do not measure time.