TT Lab
Get started
Learn Learning paths Courses

Transformers — Compute Attention By Hand

Shrink Only the Key/Value Heads

Continue in TT Lab

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

  1. 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, and the groups are divided into contiguous chunks.
  2. 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.
  3. Add expand_kv(folded, h) so that it repeats the folded heads in place and unfolds them into h heads.
  4. 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.
  5. Add heads_out(q_heads, k_heads, v_heads, g) and max_gap(left, right) so that they run the three methods with the same weights and measure the output difference.
  6. Add kv_cache_elems(layers, kv_heads, seq_len, head_dim) and kv_table(layers, h, groups, seq_len, head_dim) so that they count the number of cache elements.
  7. Add mults(n, d_model, h, g, head_dim) so that it counts the multiplications of one pass through one layer, item by item.
  8. 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

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.