Transformers — Compute Attention By Hand
Shrink Only the Key/Value Heads
Goal
Build, using only the standard library, a structure that keeps the number of query heads as it is and reduces only the number of key/value heads. You group the query heads into contiguous chunks, fold the key/value heads within each group by averaging, repeat them as many times as the number of query heads and unfold them, and then run attention. Just by changing the number of groups g, you run the three methods MHA (g=h), GQA (1with the same weights and measure the output difference, and at the end you count the number of KV cache elements and the number of multiplications separately to confirm what shrinks and what does not.
Why it matters
In today's model configurations, the number of query heads and the number of key/value heads are written separately. If the two values are equal it is MHA, if the key/value side is 1 it is MQA, and in between it is GQA. More important than memorizing the three names is knowing what shrinks.
What shrinks is KV cache memory. The number of cache elements is 2 x 층 x 키·값 헤드 수 x 길이 x 헤드 차원 (the Korean words mean layers, number of key/value heads, length and head dimension), and the number of query heads is not in this formula. On the other hand, the number of multiplications in the core of attention is h x n x n x 헤드차원 (the Korean word means head dimension), so the number of groups g does not enter at all. That is because the folded keys and values are repeated and unfolded as many times as the number of query heads and run as usual. If you count these two formulas yourself, you can answer the question "I switched to GQA, so why is prefill the same?"
This lab does not measure time. The Pod has no GPU and the CPU is shared with other work, so a speed measured here tells you nothing. Every judgment is integer counting, such as number of elements and number of multiplications, and numerical comparison with a tolerance.
Merging the key/value heads within a group by averaging is an assumption of this lab. It is the same as the method the original GQA paper uses when converting an already-trained multi-head checkpoint, but the original paper does additional training after the averaging. The output difference that comes out here does not mean quality gets worse by that much; it is a value that shows the fact that tying only the keys and values with the same weights changes the answer.
The grader does not trust the explanations you wrote down. It actually imports your module, pokes at the functions with a different number of heads and number of groups every time, and checks them against values it computes separately.
Steps
- In /root/work/tf-gqa/gqa.py, create
make_heads(count, rows, dim, seed)andgroup_of(head, h, g)andgroup_members(h, g). The samples must always give the same values for the same arguments, and the groups are divided into contiguous chunks. - Add
fold_kv(heads, g)so that it averages the key/value heads within each group position by position and folds them into g heads. - Add
expand_kv(folded, h)so that it repeats the folded heads in place and unfolds them into h heads. - Add
attend(q, k, v). Divide the scores bysqrt(헤드 차원)(the Korean word means head dimension), and for the softmax subtract the maximum before taking the exponential. - Add
heads_out(q_heads, k_heads, v_heads, g)andmax_gap(left, right)so that they run the three methods with the same weights and measure the output difference. - Add
kv_cache_elems(layers, kv_heads, seq_len, head_dim)andkv_table(layers, h, groups, seq_len, head_dim)so that they count the number of cache elements. - Add
mults(n, d_model, h, g, head_dim)so that it counts the multiplications of one pass through one layer, item by item. - Fix the samples and the model scale, make the two tables, and record the results in /root/work/tf-gqa/gqa_report.json and /root/work/tf-gqa/gqa_report.md.
Notes
- Execution contract: the grader imports
/root/work/tf-gqa/gqa.pyas a Python module and usesmake_heads,group_of,group_members,fold_kv,expand_kv,attend,heads_out,max_gap,kv_cache_elems,kv_tableandmultsdirectly. It does not run it as a script, soif __name__ == "__main__"is not needed. make_heads(count, rows, dim, seed)returns count heads, each head has rows rows, and one row is dim real numbers. The same arguments must always give the same values (therandommodule is forbidden), and a different seed must give different values. Keep the values at an absolute value of at most 4, and they must not all be the same.group_of(head, h, g)assigns in contiguous chunks. With h=8 and g=2, 0, 1, 2, 3 are group 0 and 4, 5, 6, 7 are group 1. Ifh % g != 0, or h or g is less than 1, or head is out of range, raise aValueError.group_members(h, g)is a list of length g, and each cell is the list of query head numbers in that group.fold_kv(heads, g)averages position by position the heads within a group and folds them into g heads. If g equals the number of heads, the values come out unchanged (each group has one head, so the average is itself). If it cannot be divided evenly, it is aValueError.expand_kv(folded, h)repeats each head in place. Unfolding[A, B]to 4 gives[A, A, B, B], not[A, B, A, B]. If it cannot be unfolded evenly, it is aValueError.attend(q, k, v)returns a table of the same shape as q. It uses no mask — what you want to see here is only the effect of the number of key/value heads, so you compare without the causal mask.heads_out(...)folds withfold_kv, unfolds withexpand_kv, and then callsattendonce per query head. If called withg = len(q_heads), folding and unfolding do not change the values, so the result must equal ordinary multi-head.max_gap(left, right)returns the single largest absolute difference between two outputs.kv_cache_elemsis2 * layers * kv_heads * seq_len * head_dim. The number of query heads does not enter.kv_tableis[(무리 수, 원소 수), ...](the placeholders are the number of groups and the number of elements).- The keys of
mults(n, d_model, h, g, head_dim)are the sevenproj_q,proj_k,proj_v,scores,weighted,proj_outandtotal, and all the values are integers.totalis the sum of the other six. It counts only multiplications, and not additions or exponentials. - The values of the step 8 report are set as follows. The output comparison uses
sample_h = 8,sample_n = 6andsample_head_dim = 4, withmake_heads(8, 6, 4, 101),make_heads(8, 6, 4, 202)andmake_heads(8, 6, 4, 303)as Q, K and V respectively. The numbers of groups aregroups = [8, 4, 2, 1]. - The model scale is set to
layers = 32,seq_len = 4096,head_dim = 128,h = 8,d_model = 1024andn_tokens = 4096. These are example sizes for building the table, not values measured from a real model. - The keys to put in
gqa_report.json:sample_h,sample_n,sample_head_dim,groups,diff_table,layers,seq_len,head_dim,h,d_model,n_tokens,kv_table,kv_ratio,mult_table,mult_ratio,core_mults,core_same_for_all_g.diff_tableis[[무리 수, MHA 와의 최대 차이], ...](the placeholders are the number of groups and the maximum difference from MHA),kv_ratioandmult_ratioare multiples with MHA as 1 (MHA 값 / 그 무리의 값, where the Korean words mean the MHA value divided by the value of that group), andcore_multsisscores + weighted. gqa_report.mdis written in the four sections## 무엇을 쟀나## 무리를 줄이면 출력이 얼마나 달라지나## 메모리는 줄고 곱셈은 안 준다## 어디에 쓸 것인가(the Korean headings mean "What was measured", "How much the output changes when you reduce the groups", "Memory shrinks and multiplications do not" and "Where to use it").- This Pod has neither internet nor a GPU.
pip installdoes not work, and numpy exists only inside/opt/onnx-lab/bin/python, soimport numpydoes not work in the system Python.mathalone is enough. - Official documents: the original MQA paper · the original GQA paper · Attention Is All You Need · PyTorch — MultiheadAttention
- Common mistakes: assigning groups alternately, keeping only the first head when folding, concatenating the whole list when unfolding, not dividing the scores, not subtracting the maximum in the softmax, putting the number of query heads in the cache formula, and swapping
hforgin the multiplication formula and counting that the computation shrinks too.
Divide the query heads into groups
In /root/work/tf-gqa/gqa.py, create make_heads(count, rows, dim, seed) and group_of(head, h, g) and group_members(h, g). The samples must always give the same values for the same arguments (the random module is forbidden, absolute value at most 4), and the groups are divided into contiguous chunks. With h=8 and g=2, 0, 1, 2, 3 are group 0. If it cannot be divided evenly, raise a ValueError.
A single small linear congruential formula is enough for the samples — hold the state as an integer, repeat state = (a * state + c) % m, and take out state / m - 0.5, and the same seed gives the same sequence. The group number is head // (h // g). head % g is an alternating assignment and goes off when you unfold later. If you fill group_members by calling group_of h times, the rules in the two places cannot part ways.
Fold the keys and values within a group into one
Add fold_kv(heads, g). It averages position by position the key/value heads within a group and folds them into g heads. If the number of heads is not divisible by g, it is a ValueError. If g equals the number of heads, the values must come out unchanged.
The groups use the same rule as step 1, that is, contiguous chunks. Take one group with heads[start:start + size], add position by position, and divide by the group size. If you keep only the first head, or ignore the groups and average everything, it looks the same only at g=1 by coincidence. Averaging is an assumption of this lab — it is the same as the method the original GQA paper uses when converting a checkpoint, but the original paper does additional training after that.
Unfold again to the number of query heads
Add expand_kv(folded, h). It repeats the folded heads in place and unfolds them into h heads. Unfolding [A, B] to 4 gives [A, A, B, B]. If it cannot be unfolded evenly, it is a ValueError.
The outer loop is over the folded heads and the inner loop runs h // len(folded) times. If you multiply the whole list and concatenate it (folded * size), it becomes [A, B, A, B] and query head i ends up looking at the keys and values of another group. Check that it pairs up with group_of from step 1 — what query head i looks at must be folded[group_of(i, h, g)]. If you copy the rows into new lists before returning, there is no case where fixing one place later changes several heads together.
Single-head attention
Add attend(q, k, v). Divide the scores by sqrt(헤드 차원) (the Korean word means head dimension), and for the softmax subtract the maximum before taking the exponential. The return value is a table of the same shape as q. It uses no mask.
For each query row, take the dot product with all the key rows, divide, pass through the softmax, and take the weighted average of the value rows. If you do not divide by sqrt, the softmax piles up on one position as the dimension grows. If you do not subtract the maximum, math.exp overflows on samples with large scores and raises an OverflowError — the grader deliberately puts in large values.
The three methods with the same weights
Add heads_out(q_heads, k_heads, v_heads, g) and max_gap(left, right). heads_out folds the keys and values into g heads and unfolds them to h heads, and then calls attend once per query head. If g = len(q_heads), it must not differ from ordinary multi-head in a single place, and if you reduce g, the output must differ. max_gap returns the largest absolute difference between two outputs.
It is done in three lines — fold and unfold the keys, fold and unfold the values, and attend per query head. The whole point of this step is that the query heads are not touched. max_gap sweeps all heads, rows and columns and keeps only the single largest difference. If you call it while changing g, you can see that the difference is exactly 0 at g=h and grows as g gets smaller.
Count the cache elements
Add kv_cache_elems(layers, kv_heads, seq_len, head_dim) and kv_table(layers, h, groups, seq_len, head_dim). The number of elements is 2 * layers * kv_heads * seq_len * head_dim, and the number of query heads does not enter. kv_table returns [(무리 수, 원소 수), ...] (the placeholders are the number of groups and the number of elements), and if there is a number of groups that cannot divide evenly, it is a ValueError.
The 2 at the front is the two sets, K and V. There is a spot where you feel like multiplying by the number of query heads, but queries are not kept in the cache, so they are not in the formula. In kv_table, you can just use the fact that the number of key/value heads for that number of groups is g. The reason to count in number of elements rather than bytes is that, depending on the data type, one element may be 2 bytes or 4 bytes.
Count the multiplications
Add mults(n, d_model, h, g, head_dim). The keys are proj_q, proj_k, proj_v, scores, weighted, proj_out and total, and the values are integers. scores and weighted are h * n * n * head_dim, the projections are of the form n * d_model * (헤드 수 * head_dim) (the Korean words mean number of heads), and total is the sum of the other six. If it cannot be divided evenly, it is a ValueError.
Which terms contain g and which do not is the whole of this step. After the keys and values are unfolded, attention runs as many times as the number of query heads, so g is not in scores and weighted. Only the K projection and the V projection depend on g. You count only multiplications, not additions or the exponential of the softmax — what you want to see is not the total but which terms shrink.
Record what shrinks and what does not
For the samples, use make_heads(8, 6, 4, 101), make_heads(8, 6, 4, 202) and make_heads(8, 6, 4, 303) as Q, K and V, and set the numbers of groups to [8, 4, 2, 1]. The model scale is layers = 32, seq_len = 4096, head_dim = 128, h = 8, d_model = 1024 and n_tokens = 4096. Write sample_h, sample_n, sample_head_dim, groups, diff_table, layers, seq_len, head_dim, h, d_model, n_tokens, kv_table, kv_ratio, mult_table, mult_ratio, core_mults and core_same_for_all_g in /root/work/tf-gqa/gqa_report.json, and write /root/work/tf-gqa/gqa_report.md in the four sections ## 무엇을 쟀나 ## 무리를 줄이면 출력이 얼마나 달라지나 ## 메모리는 줄고 곱셈은 안 준다 ## 어디에 쓸 것인가 (the Korean headings mean "What was measured", "How much the output changes when you reduce the groups", "Memory shrinks and multiplications do not" and "Where to use it").
Do not write the numbers by hand; fill them in with values obtained by actually running your own functions. diff_table is max_gap(그 무리의 출력, g=h 의 출력) for each number of groups (the placeholders are the output of that group and the output at g=h), so the first cell is 0.0. kv_ratio and mult_ratio are multiples with MHA as 1, so they are MHA 값 / 그 무리의 값 (the placeholders are the MHA value and the value of that group). core_mults is scores + weighted and is the same value even when the number of groups changes, so core_same_for_all_g becomes true. In the md, write the cache multiple and the multiplication multiple side by side — the conclusion of this lab is that one side shrinks by up to 8 times and the other stays almost the same.