code wiki / (root) / nx_gqa_head_map_test.nx

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}