nx_f32_attn_multi_test.nx source
↩ module page · 90 lines · 3530 B
1// nx_f32_attn_multi_test.nx -- smoke for nx_f32_attn_multi.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_attn_multi.nx"
9
10func _ulp_diff_pos(a: i64, b: i64) -> i64 {
11 if a >= b { return a - b }
12 return b - a
13}
14
15func main() -> i64 {
16 var vi: nx_int = 0
17 while vi < NX_F32_AM_N_VERDICTS {
18 if nx_f32_am_verdict_is_valid(vi) != 1 { return 5 + vi }
19 vi = vi + 1
20 }
21
22 // ===== Test A: n_tokens_q=1, n_tokens_k=1, head_dim=2 =====
23 // Q=[1, 0], K=[1, 0], V=[1, 0]
24 // score = dot(Q, K) * scale = 1.0 * scale
25 // softmax([anything]) for single element = [1.0]
26 // out = 1.0 * V = [1, 0]
27 let Q1: *i64 = sys_mmap(2 * 8) as *i64
28 let K1: *i64 = sys_mmap(2 * 8) as *i64
29 let V1: *i64 = sys_mmap(2 * 8) as *i64
30 let O1: *i64 = sys_mmap(2 * 8) as *i64
31 Q1[0] = 0x3F800000; Q1[1] = 0
32 K1[0] = 0x3F800000; K1[1] = 0
33 V1[0] = 0x3F800000; V1[1] = 0
34 let scale: i64 = 0x3F3504F3 // 1/sqrt(2) ~ 0.7071
35 let vA: nx_int = nx_f32_attn_multi(Q1, K1, V1, 1, 1, 2, 1, scale, O1)
36 if vA != NX_F32_AM_OK { return 10 + vA }
37 // Single-element softmax gives prob=1.0, so out = V
38 if _ulp_diff_pos(O1[0], V1[0]) > 4096 { return 11 }
39 if O1[1] != 0 { return 12 }
40
41 // ===== Test B: n_tokens=2, causal mask =====
42 // Q = K = V = [[1,0], [0,1]] (2 tokens, head_dim=2, orthonormal)
43 // With causal:
44 // token 0 sees only key 0: probs=[1,0], out[0] = V[0] = [1,0]
45 // token 1 sees both keys:
46 // score[0] = dot(Q[1], K[0]) * scale = 0 * scale = 0
47 // score[1] = dot(Q[1], K[1]) * scale = 1 * scale = scale
48 // probs = softmax([0, scale])
49 // out[1] = probs[0]*V[0] + probs[1]*V[1]
50 // Since V[0]=[1,0] and V[1]=[0,1], out[1] = [probs[0], probs[1]]
51 let Q2: *i64 = sys_mmap(4 * 8) as *i64
52 let K2: *i64 = sys_mmap(4 * 8) as *i64
53 let V2: *i64 = sys_mmap(4 * 8) as *i64
54 let O2: *i64 = sys_mmap(4 * 8) as *i64
55 Q2[0] = 0x3F800000; Q2[1] = 0
56 Q2[2] = 0; Q2[3] = 0x3F800000
57 K2[0] = 0x3F800000; K2[1] = 0
58 K2[2] = 0; K2[3] = 0x3F800000
59 V2[0] = 0x3F800000; V2[1] = 0
60 V2[2] = 0; V2[3] = 0x3F800000
61 let vB: nx_int = nx_f32_attn_multi(Q2, K2, V2, 2, 2, 2, 1, scale, O2)
62 if vB != NX_F32_AM_OK { return 20 + vB }
63 // token 0: out=[1,0] (only attends to key 0)
64 if _ulp_diff_pos(O2[0], 0x3F800000) > 4096 { return 21 }
65 if O2[1] != 0 { return 22 }
66 // token 1: out approx [p0, p1] where p0+p1=1, p1 > p0 since
67 // score[1] > score[0]. We just check they're both positive and
68 // sum near 1.0.
69 let sum_o2: i64 = nx_f32_add(O2[2], O2[3])
70 if _ulp_diff_pos(sum_o2, 0x3F800000) > 4096 { return 23 }
71 // p1 > p0 since score[1] > score[0]
72 if O2[3] <= O2[2] { return 24 }
73
74 // ===== Test C: non-causal full attention =====
75 // Same Q,K,V; token 0 sees both keys now.
76 let O3: *i64 = sys_mmap(4 * 8) as *i64
77 let vC: nx_int = nx_f32_attn_multi(Q2, K2, V2, 2, 2, 2, 0, scale, O3)
78 if vC != NX_F32_AM_OK { return 30 }
79 // token 0: out=[p0, p1] with p0 > p1 (since Q[0] aligned with K[0])
80 if O3[0] <= O3[1] { return 31 }
81 // Probs sum to ~1
82 let sum_o3_0: i64 = nx_f32_add(O3[0], O3[1])
83 if _ulp_diff_pos(sum_o3_0, 0x3F800000) > 4096 { return 32 }
84
85 // ===== Test D: dim error =====
86 let vD: nx_int = nx_f32_attn_multi(Q1, K1, V1, 5, 1, 2, 1, scale, O1)
87 if vD != NX_F32_AM_ERR_BAD_DIM { return 40 }
88
89 return 0
90}