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}