nx_gqa_head_map_test.nx source
↩ module page · 132 lines · 6362 B
1// nx_gqa_head_map_test.nx -- smoke for H6 GQA/MQA head map.
2
3import "nx_syscalls.nx"
4import "nx_gqa_head_map.nx"
5
6func main() -> i64 {
7 // ----- 1. MHA case (n_q == n_kv) -----
8 let mha: *NxGqaHeadMap = nx_gqa_head_map_new(32, 32)
9 if nx_gqa_head_map_is_valid(mha) != 1 { return 1 }
10 if nx_gqa_head_map_n_q_heads(mha) != 32 { return 2 }
11 if nx_gqa_head_map_n_kv_heads(mha) != 32 { return 3 }
12 if nx_gqa_head_map_group_size(mha) != 1 { return 4 }
13 if nx_gqa_head_map_kind(mha) != NX_ATTN_MHA { return 5 }
14 // Each Q head maps to a unique KV head equal to its own index.
15 var i: i64 = 0
16 while i < 32 {
17 if nx_gqa_head_map_kv_head_for_q_head(mha, i) != i { return 6 }
18 i = i + 1
19 }
20 // No KV savings under MHA.
21 if nx_gqa_head_map_kv_savings_q16(mha) != 0 { return 7 }
22
23 // ----- 2. GQA case (Llama 3 70B style: 64 Q heads, 8 KV heads, group 8) -----
24 let gqa: *NxGqaHeadMap = nx_gqa_head_map_new(64, 8)
25 if nx_gqa_head_map_is_valid(gqa) != 1 { return 8 }
26 if nx_gqa_head_map_group_size(gqa) != 8 { return 9 }
27 if nx_gqa_head_map_kind(gqa) != NX_ATTN_GQA { return 10 }
28 // Q heads 0..7 -> KV 0; 8..15 -> KV 1; ... 56..63 -> KV 7.
29 if nx_gqa_head_map_kv_head_for_q_head(gqa, 0) != 0 { return 11 }
30 if nx_gqa_head_map_kv_head_for_q_head(gqa, 7) != 0 { return 12 }
31 if nx_gqa_head_map_kv_head_for_q_head(gqa, 8) != 1 { return 13 }
32 if nx_gqa_head_map_kv_head_for_q_head(gqa, 15) != 1 { return 14 }
33 if nx_gqa_head_map_kv_head_for_q_head(gqa, 56) != 7 { return 15 }
34 if nx_gqa_head_map_kv_head_for_q_head(gqa, 63) != 7 { return 16 }
35 // Group occupancy: 8 Q heads per KV head, for all 8 KV heads.
36 var kv: i64 = 0
37 while kv < 8 {
38 if nx_gqa_head_map_count_q_heads_in_group(gqa, kv) != 8 { return 17 }
39 kv = kv + 1
40 }
41 // KV savings = (64 - 8) / 64 = 0.875 in Q16 = 0.875 * 65536 = 57344
42 if nx_gqa_head_map_kv_savings_q16(gqa) != 57344 { return 18 }
43
44 // ----- 3. MQA case (n_kv == 1) -----
45 let mqa: *NxGqaHeadMap = nx_gqa_head_map_new(32, 1)
46 if nx_gqa_head_map_is_valid(mqa) != 1 { return 19 }
47 if nx_gqa_head_map_group_size(mqa) != 32 { return 20 }
48 if nx_gqa_head_map_kind(mqa) != NX_ATTN_MQA { return 21 }
49 // ALL Q heads -> KV head 0.
50 var j: i64 = 0
51 while j < 32 {
52 if nx_gqa_head_map_kv_head_for_q_head(mqa, j) != 0 { return 22 }
53 j = j + 1
54 }
55 if nx_gqa_head_map_count_q_heads_in_group(mqa, 0) != 32 { return 23 }
56 // KV savings = (32 - 1) / 32 = 0.96875 in Q16 = 0.96875 * 65536 = 63488
57 if nx_gqa_head_map_kv_savings_q16(mqa) != 63488 { return 24 }
58
59 // ----- 4. Llama 2 7B style (32 Q, 32 KV) is MHA -----
60 let l27b: *NxGqaHeadMap = nx_gqa_head_map_new(32, 32)
61 if nx_gqa_head_map_kind(l27b) != NX_ATTN_MHA { return 25 }
62
63 // Llama 3 8B style (32 Q, 8 KV) is GQA group 4.
64 let l38b: *NxGqaHeadMap = nx_gqa_head_map_new(32, 8)
65 if nx_gqa_head_map_kind(l38b) != NX_ATTN_GQA { return 26 }
66 if nx_gqa_head_map_group_size(l38b) != 4 { return 27 }
67
68 // ----- 5. Divisibility rejected -----
69 if (nx_gqa_head_map_new(32, 7) as i64) != 0 { return 28 } // 32 % 7 != 0
70 if (nx_gqa_head_map_new(64, 6) as i64) != 0 { return 29 } // 64 % 6 != 0
71 if (nx_gqa_head_map_new(64, 5) as i64) != 0 { return 30 } // 64 % 5 != 0
72
73 // ----- 6. n_kv > n_q rejected -----
74 if (nx_gqa_head_map_new(8, 16) as i64) != 0 { return 31 }
75
76 // ----- 7. Bad inputs -----
77 if (nx_gqa_head_map_new(0, 1) as i64) != 0 { return 32 }
78 if (nx_gqa_head_map_new(-1, 1) as i64) != 0 { return 33 }
79 if (nx_gqa_head_map_new(32, 0) as i64) != 0 { return 34 }
80 if (nx_gqa_head_map_new(32, -1) as i64) != 0 { return 35 }
81 if (nx_gqa_head_map_new(NX_GQA_MAX_Q_HEADS + 1, 1) as i64) != 0 { return 36 }
82
83 // ----- 8. Out-of-range lookups -----
84 if nx_gqa_head_map_kv_head_for_q_head(gqa, -1) != (0 - NX_GQA_OUT_OF_RANGE) { return 37 }
85 if nx_gqa_head_map_kv_head_for_q_head(gqa, 64) != (0 - NX_GQA_OUT_OF_RANGE) { return 38 }
86 if nx_gqa_head_map_kv_head_for_q_head(gqa, 100) != (0 - NX_GQA_OUT_OF_RANGE) { return 39 }
87 if nx_gqa_head_map_count_q_heads_in_group(gqa, -1) != (0 - NX_GQA_OUT_OF_RANGE) { return 40 }
88 if nx_gqa_head_map_count_q_heads_in_group(gqa, 8) != (0 - NX_GQA_OUT_OF_RANGE) { return 41 }
89
90 // ----- 9. Tamper -----
91 let tamper_m: *NxGqaHeadMap = nx_gqa_head_map_new(32, 8)
92 tamper_m.canary_post = 0xDEADBEEF
93 if nx_gqa_head_map_is_valid(tamper_m) != 0 { return 42 }
94 if nx_gqa_head_map_kv_head_for_q_head(tamper_m, 0) != (0 - NX_GQA_TAMPER) { return 43 }
95 if nx_gqa_head_map_count_q_heads_in_group(tamper_m, 0) != (0 - NX_GQA_TAMPER) { return 44 }
96 if nx_gqa_head_map_n_q_heads(tamper_m) != -1 { return 45 }
97 if nx_gqa_head_map_kind(tamper_m) != -1 { return 46 }
98 if nx_gqa_head_map_kv_savings_q16(tamper_m) != -1 { return 47 }
99
100 // ----- 10. Classify() pure helper -----
101 if nx_attn_classify(32, 32) != NX_ATTN_MHA { return 48 }
102 if nx_attn_classify(32, 8) != NX_ATTN_GQA { return 49 }
103 if nx_attn_classify(32, 1) != NX_ATTN_MQA { return 50 }
104 if nx_attn_classify(64, 8) != NX_ATTN_GQA { return 51 }
105
106 // ----- 11. Sealed-enum gates -----
107 if nx_attn_kind_is_valid(NX_ATTN_MHA) != 1 { return 52 }
108 if nx_attn_kind_is_valid(NX_ATTN_GQA) != 1 { return 53 }
109 if nx_attn_kind_is_valid(NX_ATTN_MQA) != 1 { return 54 }
110 if nx_attn_kind_is_valid(-1) != 0 { return 55 }
111 if nx_attn_kind_is_valid(NX_ATTN_N_KINDS) != 0 { return 56 }
112 if nx_gqa_verdict_is_valid(NX_GQA_OK) != 1 { return 57 }
113 if nx_gqa_verdict_is_valid(NX_GQA_TAMPER) != 1 { return 58 }
114 if nx_gqa_verdict_is_valid(-1) != 0 { return 59 }
115 if nx_gqa_verdict_is_valid(NX_GQA_N_VERDICTS) != 0 { return 60 }
116
117 // ----- 12. group_size * n_kv == n_q invariant -----
118 if l38b.group_size * l38b.n_kv_heads != l38b.n_q_heads { return 61 }
119 if gqa.group_size * gqa.n_kv_heads != gqa.n_q_heads { return 62 }
120 if mqa.group_size * mqa.n_kv_heads != mqa.n_q_heads { return 63 }
121
122 // ----- 13. Sum of group counts == n_q_heads invariant -----
123 var sum_counts: i64 = 0
124 var kvi: i64 = 0
125 while kvi < gqa.n_kv_heads {
126 sum_counts = sum_counts + nx_gqa_head_map_count_q_heads_in_group(gqa, kvi)
127 kvi = kvi + 1
128 }
129 if sum_counts != gqa.n_q_heads { return 64 }
130
131 return 0
132}