nx_f32_softmax_test.nx source
↩ module page · 74 lines · 2937 B
1// nx_f32_softmax_test.nx -- smoke for nx_f32_softmax.nx.
2//
3// Stable softmax KATs. v1 polynomial exp has ~1000 ULP error;
4// downstream softmax probabilities tolerate ~5e-3 relative error
5// for these canonical cases.
6
7import "nx_syscalls.nx"
8import "nx_tier.nx"
9import "nx_f32.nx"
10import "nx_f32_div.nx"
11import "nx_f32_exp.nx"
12import "nx_f32_softmax.nx"
13
14func _ulp_diff_pos(a: i64, b: i64) -> i64 {
15 if a >= b { return a - b }
16 return b - a
17}
18
19func main() -> i64 {
20 // Verdict gate
21 var vi: nx_int = 0
22 while vi < NX_F32_SM_N_VERDICTS {
23 if nx_f32_sm_verdict_is_valid(vi) != 1 { return 5 + vi }
24 vi = vi + 1
25 }
26
27 // ===== Test 1: uniform x=[0,0,0] -> out=[1/3, 1/3, 1/3] =====
28 let x1: *i64 = sys_mmap(3 * 8) as *i64
29 let o1: *i64 = sys_mmap(3 * 8) as *i64
30 x1[0] = 0; x1[1] = 0; x1[2] = 0
31 let v1: nx_int = nx_f32_softmax(x1, 3, o1)
32 if v1 != NX_F32_SM_OK { return 10 }
33 // 1/3 in f32 = 0x3EAAAAAB (the nearest IEEE 754 representation)
34 if _ulp_diff_pos(o1[0], 0x3EAAAAAB) > 4096 { return 11 }
35 if _ulp_diff_pos(o1[1], 0x3EAAAAAB) > 4096 { return 12 }
36 if _ulp_diff_pos(o1[2], 0x3EAAAAAB) > 4096 { return 13 }
37
38 // ===== Test 2: probabilities sum to 1.0 =====
39 let sum_check: i64 = nx_f32_add(nx_f32_add(o1[0], o1[1]), o1[2])
40 // Sum should be very close to 1.0 = 0x3F800000.
41 if _ulp_diff_pos(sum_check, 0x3F800000) > 4096 { return 20 }
42
43 // ===== Test 3: peaked input [10, 0, 0] -> heavily concentrated =====
44 let x3: *i64 = sys_mmap(3 * 8) as *i64
45 let o3: *i64 = sys_mmap(3 * 8) as *i64
46 x3[0] = 0x41200000; x3[1] = 0; x3[2] = 0 // 10, 0, 0
47 let v3: nx_int = nx_f32_softmax(x3, 3, o3)
48 if v3 != NX_F32_SM_OK { return 30 }
49 // out[0] should be very close to 1.0; out[1]=out[2] very small.
50 // exp(10) / (exp(10) + 2) ~= 22026 / 22028 ~= 0.99991
51 // 0.99991 in f32 ~= 0x3F7FF0BD or similar
52 // Just check out[0] > 0.999 = 0x3F7FBE77
53 if (o3[0] >> 31) & 1 != 0 { return 31 } // positive
54 let abs_o3_0: i64 = o3[0] & 0x7FFFFFFF
55 if abs_o3_0 < 0x3F7FBE77 { return 32 } // > 0.999
56
57 // ===== Test 4: max-subtraction prevents overflow =====
58 // x = [100, 100, 100]. Without max-subtraction, exp(100)
59 // would overflow. With it, shifted = [0,0,0], output uniform.
60 let x4: *i64 = sys_mmap(3 * 8) as *i64
61 let o4: *i64 = sys_mmap(3 * 8) as *i64
62 x4[0] = 0x42C80000; x4[1] = 0x42C80000; x4[2] = 0x42C80000 // 100s
63 let v4: nx_int = nx_f32_softmax(x4, 3, o4)
64 if v4 != NX_F32_SM_OK { return 40 }
65 if _ulp_diff_pos(o4[0], 0x3EAAAAAB) > 4096 { return 41 }
66 if _ulp_diff_pos(o4[1], 0x3EAAAAAB) > 4096 { return 42 }
67 if _ulp_diff_pos(o4[2], 0x3EAAAAAB) > 4096 { return 43 }
68
69 // ===== Test 5: probabilities sum to 1 in stressful case =====
70 let sum4: i64 = nx_f32_add(nx_f32_add(o4[0], o4[1]), o4[2])
71 if _ulp_diff_pos(sum4, 0x3F800000) > 4096 { return 50 }
72
73 return 0
74}