TT Lab
Get started
Learn Learning paths Courses

Transformers — Compute Attention By Hand

Count the Cost of Context Length Yourself

Continue in TT Lab

Goal

Build a counter that counts multiplications, and count directly what grows as n squared and what grows in proportion to n as the context length n increases. Build a table that puts the two shares side by side, and find the crossover point where the n-squared share catches up with the linear share, using your own counted numbers. You also count in integers that the scores left by the causal mask number only n(n+1)/2, that holding the whole score matrix means n squared elements, and that attaching one token produces only one new row of scores.

Why it matters

The statement "quadratic in context length" is only half right. Two computations of different character are mixed inside one layer. The scoring and mixing of attention sweep the whole context at every position and so grow as a square, but the query/key/value projections, the output projection and the feed-forward are work that runs once for each position in the same way, so they are proportional only to the number of positions. Which side dominates is decided by the length and the model width together. So if you estimate the cost of extending the context limit with a single ratio, you are sure to be wrong. This lab does not measure time. The same code differs with the machine and the load and wobbles even more inside a container, so there is nothing to compare by measuring. Instead it counts the number of multiplications. The count is the same integer wherever you run it, so you can show it to others as evidence. Every number that appears here is a value your counter counted, and seconds or GB of a real model are not measured, so they are not used. For the model shape, it uses the base setting of Attention Is All You Need as an assumption — D_MODEL = 512, D_FF = 2048, N_HEADS = 8. This is an assumption fixed by this lab, not the values of the model you use. The grader does not trust the explanations you wrote down. It actually imports your module, pokes at the functions with different sizes every time, and even checks how far the counter actually went up. The sizes change on every run, so you cannot memorize values and plug them in.

Steps

  1. In /root/work/tf-cost/cost.py, create the constants D_MODEL = 512, D_FF = 2048, N_HEADS = 8, a counter MulCount, and dot(a, b, ctr), which uses it.
  2. Add attn_scores(Q, K, ctr), attn_mix(A, V, ctr) and quad_mults(n, d_model) so that they count the share that grows as n squared.
  3. Add matvec(M, x, ctr), per_position_mults(d_model, d_ff, ctr) and linear_mults(n, d_model, d_ff) so that they count the share proportional to the number of positions.
  4. Create cost_table(ns, d_model, d_ff) so that it returns (n, n제곱 몫, 선형 몫, 합) for each length (the placeholders are the length, the n-squared share, the linear share and the sum).
  5. Create crossover_n(d_model, d_ff) so that it finds the length at which the n-squared share first catches up with the linear share.
  6. Create causal_pairs(n) and wasted_pairs(n) so that they count the scores the causal mask leaves and the cells it throws away.
  7. Create score_bytes(n, n_heads, itemsize), max_context_for_bytes(budget_bytes, n_heads, itemsize), append_pairs(n) and generate_pairs(n, g).
  8. Record the results of actually running the functions above in /root/work/tf-cost/cost_report.json and /root/work/tf-cost/cost_report.md.

Notes

What is not measured

Time is not measured. The GB or seconds of a real model are not used either. If you write a number you have not measured in the record, that record has no basis.

Build a yardstick that counts multiplications

In /root/work/tf-cost/cost.py, create the constants D_MODEL = 512, D_FF = 2048, N_HEADS = 8, a counter MulCount (it holds mults and calls and is raised with add(k)), and dot(a, b, ctr). dot returns the dot product value and raises the counter by exactly the vector length.

Do not try to measure time — the same code differs with the machine and the load and cannot be compared. A dot product of length d is exactly d multiplications, so the single line ctr.add(len(a)) is enough. If you count one call as 1, every later number collapses. calls counts how many separate times it was written to the ledger, and is used in step 3.

The share that grows as n squared

Add attn_scores(Q, K, ctr), attn_mix(A, V, ctr) and quad_mults(n, d_model). attn_scores returns the n x n score matrix and attn_mix returns the n x d result of mixing V with A, and both raise the counter by n * n * d. quad_mults is the two combined, 2 * n * n * d_model.

attn_scores is done by taking dot of each query with all the keys. attn_mix is where people get confused — what comes out is small, n x d, but at each output position it has to sweep the whole context of n entries, so the multiplications number n * n * d, exactly the same as the score computation. If you pull out a column of V and pass it to dot, the counter comes out right on its own. The number of heads does not enter quad_mults.

The share proportional only to the number of positions

Add matvec(M, x, ctr), per_position_mults(d_model, d_ff, ctr) and linear_mults(n, d_model, d_ff). per_position_mults counts the multiplications by actually running each of the six linear maps (four projections, two feed-forward) and returns how much it raised the counter. linear_mults is the integer you get by multiplying that value by the number of positions.

The values of the matrices do not matter — counting is the purpose, so you can just build and run matrices with the right shapes. The six are four of d_model x d_model, one of d_ff x d_model and one of d_model x d_ff. The vector you pass to feed-forward layer 2 is the one of length d_ff that layer 1 returned. The grader also checks that the counter was called at least six times, so if you add one formula in a single step and finish, you fail. linear_mults does not take a counter.

Put the two shares side by side

Create cost_table(ns, d_model, d_ff). For each length in ns, return the four cells (n, quad, linear, quad + linear), in the same order as ns.

If you use the quad_mults and linear_mults you built earlier as they are, it is five lines. Read it while doubling the length each time — the first cell goes up fourfold each time, and the later cell goes up twofold each time. The last cell must be the sum of the first two without fail. If you write down only one side, it goes off when you look for the crossover point later.

Find the crossover point

Create crossover_n(d_model, d_ff). It is the n at which quad_mults(n, d_model) >= linear_mults(n, d_model, d_ff) first becomes true. It includes equality.

You may find it by going up from 1 or solve it by hand. If you divide both sides by n and 2 * d_model, the condition becomes very short. If you leave out the equality and use only the inequality, the answer is off by exactly one. Try calling it while varying the model width — you can see in numbers that the wider the width, the further the crossover point shifts back.

Count the half the mask throws away

Create causal_pairs(n) and wasted_pairs(n). causal_pairs is the number of scores actually used under the causal mask, and wasted_pairs is n * n minus that share.

The i-th query looks only up to itself. The row length grows 1, 2, 3 and ends at n, so if you count, it is n * (n + 1) / 2. You must not drop the diagonal — looking at itself is correct. causal_pairs(0) is 0 and causal_pairs(1) is 1. Print a few and see where the discarded fraction approaches as the length grows.

Memory and scores that grow one row at a time

Create score_bytes(n, n_heads, itemsize), max_context_for_bytes(budget_bytes, n_heads, itemsize), append_pairs(n) and generate_pairs(n, g). The first two are the bytes when holding the whole score matrix and the longest context that fits within a budget, and the last two are the number of scores newly produced when attaching one token and the total while producing g tokens.

It is easy to leave out the number of heads in score_bytes. In max_context_for_bytes, rounding a real-number square root can overshoot by one, so use math.isqrt — divide the budget by heads and bytes and then take the integer square root. append_pairs(n) is n + 1, because the one new query looks at n + 1 entries including itself. generate_pairs is the sum from n+1 to n+g, and be sure to check that generate_pairs(0, N) equals causal_pairs(N).

Leave the counted numbers as a record

Actually run the functions above and write d_model, d_ff, n_heads, itemsize, table, quad_ratio, linear_ratio, crossover_n, quad_at_crossover, linear_at_crossover, causal_n, causal_pairs, wasted_pairs, wasted_fraction, score_bytes_at_causal_n, max_context_1gib, append_pairs_at_causal_n, generate_pairs_at_causal_n and identity_ok in /root/work/tf-cost/cost_report.json, and record /root/work/tf-cost/cost_report.md in the five sections ## 무엇을 세었나 ## 두 배로 늘리면 무엇이 네 배가 되나 ## 교차점은 어디인가 ## 인과 마스크가 버리는 절반 ## 메모리와 한 토큰씩 늘어나는 점수 (the Korean headings mean "What was counted", "What becomes four times when you double", "Where the crossover point is", "The half the causal mask throws away" and "Memory and scores that grow one token at a time").

Do not write the numbers by hand; fill them in with values obtained by running your own code. The table lengths are [128, 256, 512, 1024, 2048, 4096], and quad_ratio and linear_ratio are the ratios between the last two rows. causal_n is 2048, the data type is 2 bytes, the budget is 1073741824, and the number of generated tokens is 256. identity_ok is whether generate_pairs(0, causal_n) == causal_pairs(causal_n) is true. In the record, do not write numbers you have not measured — seconds or GB of a real model have never been measured here.