TT Lab
Get started
Learn Learning paths Courses

Transformers — Compute Attention By Hand

Count the Multiplications a KV Cache Saves

Continue in TT Lab

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

  1. In /root/work/tf-kv/kv.py, create CONFIG and reset_muls(), muls(), mul(a, b) and dot(u, v). mul multiplies and raises the count by one, and dot multiplies only through that mul.
  2. Add project(x, W) so that it passes one token through one set of weights. W is a matrix whose rows each have length len(x), and the length of the list returned is len(W).
  3. 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.
  4. 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.
  5. Add new_cache(), store(cache, k, v, digits=None) and step_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.
  6. Add ATOL, RTOL, close(got, want) and compare(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.
  7. Add mul_table(xs, Wq, Wk, Wv) so that it counts the multiplications of the two methods for context lengths from 1 to len(xs). The return value is a list of (문맥길이, 캐시없음, 캐시있음) tuples (the placeholders are the context length, without a cache and with a cache).
  8. 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

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.