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}