code wiki / (root) / nx_f32_mha_test.nx

nx_f32_mha_test.nx source

↩ module page · 92 lines · 3229 B

1// nx_f32_mha_test.nx -- smoke for nx_f32_mha.nx. 2 3import "nx_syscalls.nx" 4import "nx_tier.nx" 5import "nx_f32.nx" 6import "nx_f32_matmul.nx" 7import "nx_f32_softmax.nx" 8import "nx_f32_mha.nx" 9 10func main() -> i64 { 11 var vi: nx_int = 0 12 while vi < NX_F32_MHA_N_VERDICTS { 13 if nx_f32_mha_verdict_is_valid(vi) != 1 { return 5 + vi } 14 vi = vi + 1 15 } 16 17 // ===== Test A: 2-head MHA, no GQA, zero weights -> zero output ===== 18 // hidden_dim=4, n_heads=2, n_kv_heads=2, head_dim=2 19 let hidden_dim: nx_int = 4 20 let n_heads: nx_int = 2 21 let n_kv_heads: nx_int = 2 22 let head_dim: nx_int = 2 23 24 let x: *i64 = sys_mmap(hidden_dim * 8) as *i64 25 x[0] = 0x3F800000 // 1.0 26 x[1] = 0x40000000 // 2.0 27 x[2] = 0x40400000 // 3.0 28 x[3] = 0x40800000 // 4.0 29 30 // Zero weight matrices (mmap returns zeroed memory). 31 let W_q: *i64 = sys_mmap(hidden_dim * hidden_dim * 8) as *i64 32 let W_k: *i64 = sys_mmap(hidden_dim * hidden_dim * 8) as *i64 33 let W_v: *i64 = sys_mmap(hidden_dim * hidden_dim * 8) as *i64 34 let W_o: *i64 = sys_mmap(hidden_dim * hidden_dim * 8) as *i64 35 36 let out: *i64 = sys_mmap(hidden_dim * 8) as *i64 37 38 // attn_scale = 1/sqrt(2) ~= 0x3F3504F3 39 let attn_scale: i64 = 0x3F3504F3 40 41 let vA: nx_int = nx_f32_mha_single_token( 42 x, hidden_dim, n_heads, n_kv_heads, head_dim, 43 W_q, W_k, W_v, W_o, attn_scale, out) 44 if vA != NX_F32_MHA_OK { return 10 + vA } 45 46 // With zero weights, every matmul produces 0; out should be all zero. 47 if out[0] != 0 { return 20 } 48 if out[1] != 0 { return 21 } 49 if out[2] != 0 { return 22 } 50 if out[3] != 0 { return 23 } 51 52 // ===== Test B: GQA with n_heads=4, n_kv_heads=2 ===== 53 // hidden_dim=8, n_heads=4, n_kv_heads=2, head_dim=2. group_size=2. 54 // Still zero weights -> zero output. 55 let xB: *i64 = sys_mmap(8 * 8) as *i64 56 xB[0] = 0x3F800000; xB[1] = 0x40000000 57 xB[2] = 0x40400000; xB[3] = 0x40800000 58 xB[4] = 0x40A00000; xB[5] = 0x40C00000 59 xB[6] = 0x40E00000; xB[7] = 0x41000000 60 61 let W_q_b: *i64 = sys_mmap(8 * 8 * 8) as *i64 // [8, 8] for Q 62 let W_k_b: *i64 = sys_mmap(8 * 4 * 8) as *i64 // [8, 4] for K (kv_dim=4) 63 let W_v_b: *i64 = sys_mmap(8 * 4 * 8) as *i64 64 let W_o_b: *i64 = sys_mmap(8 * 8 * 8) as *i64 65 let out_b: *i64 = sys_mmap(8 * 8) as *i64 66 67 let vB: nx_int = nx_f32_mha_single_token( 68 xB, 8, 4, 2, 2, 69 W_q_b, W_k_b, W_v_b, W_o_b, attn_scale, out_b) 70 if vB != NX_F32_MHA_OK { return 30 + vB } 71 72 var i: nx_int = 0 73 while i < 8 { 74 if out_b[i] != 0 { return 40 + i } 75 i = i + 1 76 } 77 78 // ===== Test C: bad-dim error (hidden != n_heads * head_dim) ===== 79 let vC: nx_int = nx_f32_mha_single_token( 80 x, 5, 2, 2, 2, 81 W_q, W_k, W_v, W_o, attn_scale, out) 82 if vC != NX_F32_MHA_ERR_BAD_DIM { return 60 } 83 84 // ===== Test D: GQA-misalignment error (n_heads % n_kv_heads != 0) ===== 85 let vD: nx_int = nx_f32_mha_single_token( 86 x, 4, 3, 2, 2, // Wait n_heads*head_dim = 6 != hidden_dim=4, hits BAD_DIM first 87 W_q, W_k, W_v, W_o, attn_scale, out) 88 // Either error is acceptable; expect non-OK. 89 if vD == NX_F32_MHA_OK { return 70 } 90 91 return 0 92}