Transformers — Compute Attention By Hand
Count the Cost of Context Length Yourself
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
- In /root/work/tf-cost/cost.py, create the constants
D_MODEL = 512,D_FF = 2048,N_HEADS = 8, a counterMulCount, anddot(a, b, ctr), which uses it. - Add
attn_scores(Q, K, ctr),attn_mix(A, V, ctr)andquad_mults(n, d_model)so that they count the share that grows as n squared. - Add
matvec(M, x, ctr),per_position_mults(d_model, d_ff, ctr)andlinear_mults(n, d_model, d_ff)so that they count the share proportional to the number of positions. - 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). - 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. - Create
causal_pairs(n)andwasted_pairs(n)so that they count the scores the causal mask leaves and the cells it throws away. - Create
score_bytes(n, n_heads, itemsize),max_context_for_bytes(budget_bytes, n_heads, itemsize),append_pairs(n)andgenerate_pairs(n, g). - 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
- Execution contract: the grader imports
/root/work/tf-cost/cost.pyas a Python module and usesD_MODEL,D_FF,N_HEADS,MulCount,dot,attn_scores,attn_mix,quad_mults,matvec,per_position_mults,linear_mults,cost_table,crossover_n,causal_pairs,wasted_pairs,score_bytes,max_context_for_bytes,append_pairsandgenerate_pairsdirectly. It does not run it as a script, soif __name__ == "__main__"is not needed. MulCountholds two values,mults(the number of multiplications counted so far) andcalls(the number of times it was written to the ledger), and adds k withadd(k). Both start at 0.dot(a, b, ctr)returns the dot product value and raises the counter by exactly the length. It counts one multiplication as 1, not one call as 1.attn_scores(Q, K, ctr)returns a matrix of size (number of rows of Q) x (number of rows of K). If Q and K are each n x d, the counter goes up byn * n * d.attn_mix(A, V, ctr)returns n x d when A is n x n and V is n x d, and the counter goes up byn * n * dagain. The computation is not small just because the output is small.quad_mults(n, d_model)is the two shares combined,2 * n * n * d_model. The number of heads does not enter — because each head's width shrinks tod_model / hand there are h such heads, so the sum is the same.per_position_mults(d_model, d_ff, ctr)counts multiplications by actually running each of the six linear maps. They are the three query/key/value projections (d_model x d_model), one output projection (d_model x d_model), feed-forward layer 1 (d_ff x d_model) and layer 2 (d_model x d_ff). The return value is how much this function raised the counter, and the grader also checks that the counter was called at least six times. If you add one formula in a single step and finish, you fail.linear_mults(n, d_model, d_ff)is the integer you get by multiplying the value for one position by the number of positions. It does not take a counter.- The rows that
cost_table(ns, d_model, d_ff)returns have four cells,(n, quad, linear, quad + linear), in the same order asns. crossover_n(d_model, d_ff)is the n at whichquad_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 by solving the formula.causal_pairs(n)comes from the fact that the i-th row usesi + 1entries.causal_pairs(0)is 0 andcausal_pairs(1)is 1.wasted_pairs(n)isn * nminus the share actually used.score_bytes(n, n_heads, itemsize)isn * n * n_heads * itemsize. Do not leave out the number of heads.max_context_for_bytes(budget_bytes, n_heads, itemsize)is the largest n that satisfiesscore_bytes(n, ...) <= budget_bytes. Rounding a real-number square root can overshoot by one, so usemath.isqrt.append_pairs(n)isn + 1, andgenerate_pairs(n, g)is the sum of the row lengths fromn+1ton+g.generate_pairs(0, N)must equalcausal_pairs(N).- The step 8 report uses
D_MODEL,D_FF,N_HEADS,ITEMSIZE = 2(assuming a two-byte data type), a budget of1073741824(1 GiB), the table lengths[128, 256, 512, 1024, 2048, 4096], a reference length of2048for the causal and memory calculations, and256generated tokens. quad_ratioandlinear_ratioare the ratios between the last two rows of the table (2048 and 4096). They are divisions, so they are floating-point numbers, and the grader compares them withabs(a - b) <= atol + rtol * abs(b).- This Pod has no internet.
pip installdoes not work, and the system Python has no numpy, torch or transformers. numpy exists only inside/opt/onnx-lab/bin/python. The standard library alone is enough. - Official documents: Attention Is All You Need · PyTorch — scaled_dot_product_attention · Hugging Face — Cache strategies
- Common mistakes: counting the counter by number of calls, thinking
attn_mixcounts only as much as the output size, leaving out the output projection from the linear share, finding the crossover point without equality so that it is off by one, dropping the diagonal incausal_pairs, and not multiplyingscore_bytesby the number of heads.
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.