code wiki / (root) / nx_f32_softmax_test.nx

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}