code wiki / (root) / nx_f32_mha_multi_test.nx

nx_f32_mha_multi_test.nx source

↩ module page · 91 lines · 3250 B

1// nx_f32_mha_multi_test.nx -- smoke for nx_f32_mha_multi.nx. 2 3import "nx_syscalls.nx" 4import "nx_tier.nx" 5import "nx_f32.nx" 6import "nx_f32_matmul.nx" 7import "nx_f32_attn_multi.nx" 8import "nx_f32_mha_multi.nx" 9 10func main() -> i64 { 11 var vi: nx_int = 0 12 while vi < NX_F32_MHAM_N_VERDICTS { 13 if nx_f32_mham_verdict_is_valid(vi) != 1 { return 5 + vi } 14 vi = vi + 1 15 } 16 17 // Common params: hidden_dim=4, n_heads=2, n_kv_heads=2, head_dim=2. 18 let hidden_dim: nx_int = 4 19 let n_heads: nx_int = 2 20 let n_kv_heads: nx_int = 2 21 let head_dim: nx_int = 2 22 23 // Zero weights -> zero output regardless of input. 24 let W_q: *i64 = sys_mmap(hidden_dim * hidden_dim * 8) as *i64 25 let W_k: *i64 = sys_mmap(hidden_dim * hidden_dim * 8) as *i64 26 let W_v: *i64 = sys_mmap(hidden_dim * hidden_dim * 8) as *i64 27 let W_o: *i64 = sys_mmap(hidden_dim * hidden_dim * 8) as *i64 28 29 let attn_scale: i64 = 0x3F3504F3 // 1/sqrt(2) 30 31 // ===== Test A: n_tokens=1, zero weights -> zero output ===== 32 let xA: *i64 = sys_mmap(1 * hidden_dim * 8) as *i64 33 let outA: *i64 = sys_mmap(1 * hidden_dim * 8) as *i64 34 xA[0] = 0x3F800000; xA[1] = 0x40000000 35 xA[2] = 0x40400000; xA[3] = 0x40800000 36 37 let vA: nx_int = nx_f32_mha_multi_token( 38 xA, 1, hidden_dim, n_heads, n_kv_heads, head_dim, 39 W_q, W_k, W_v, W_o, 1, attn_scale, outA) 40 if vA != NX_F32_MHAM_OK { return 10 + vA } 41 if outA[0] != 0 { return 20 } 42 if outA[1] != 0 { return 21 } 43 if outA[2] != 0 { return 22 } 44 if outA[3] != 0 { return 23 } 45 46 // ===== Test B: n_tokens=2, causal, zero weights -> zero output ===== 47 let xB: *i64 = sys_mmap(2 * hidden_dim * 8) as *i64 48 let outB: *i64 = sys_mmap(2 * hidden_dim * 8) as *i64 49 var i: nx_int = 0 50 while i < 8 { xB[i] = 0x3F800000; i = i + 1 } // all 1.0 51 52 let vB: nx_int = nx_f32_mha_multi_token( 53 xB, 2, hidden_dim, n_heads, n_kv_heads, head_dim, 54 W_q, W_k, W_v, W_o, 1, attn_scale, outB) 55 if vB != NX_F32_MHAM_OK { return 30 + vB } 56 var j: nx_int = 0 57 while j < 8 { 58 if outB[j] != 0 { return 40 + j } 59 j = j + 1 60 } 61 62 // ===== Test C: n_tokens=3, non-causal, GQA 4-heads x 2-kv-heads ===== 63 // hidden_dim=8, head_dim=2, kv_dim=4 64 let xC: *i64 = sys_mmap(3 * 8 * 8) as *i64 65 let outC: *i64 = sys_mmap(3 * 8 * 8) as *i64 66 let W_qC: *i64 = sys_mmap(8 * 8 * 8) as *i64 // [8, 8] 67 let W_kC: *i64 = sys_mmap(8 * 4 * 8) as *i64 // [8, 4] 68 let W_vC: *i64 = sys_mmap(8 * 4 * 8) as *i64 // [8, 4] 69 let W_oC: *i64 = sys_mmap(8 * 8 * 8) as *i64 // [8, 8] 70 71 var k: nx_int = 0 72 while k < 24 { xC[k] = 0x3F800000; k = k + 1 } // input doesn't matter 73 74 let vC: nx_int = nx_f32_mha_multi_token( 75 xC, 3, 8, 4, 2, 2, 76 W_qC, W_kC, W_vC, W_oC, 0, attn_scale, outC) 77 if vC != NX_F32_MHAM_OK { return 60 + vC } 78 var m: nx_int = 0 79 while m < 24 { 80 if outC[m] != 0 { return 70 + m } 81 m = m + 1 82 } 83 84 // ===== Test D: bad-dim error ===== 85 let vD: nx_int = nx_f32_mha_multi_token( 86 xA, 1, 5, n_heads, n_kv_heads, head_dim, 87 W_q, W_k, W_v, W_o, 1, attn_scale, outA) 88 if vD != NX_F32_MHAM_ERR_BAD_DIM { return 100 } 89 90 return 0 91}