code wiki / (root) / nx_f32_attn_multi_test.nx

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}