code wiki / (root) / nx_f32_rmsnorm_test.nx

nx_f32_rmsnorm_test.nx source

↩ module page · 87 lines · 3203 B

1// nx_f32_rmsnorm_test.nx -- smoke for nx_f32_rmsnorm.nx. 2 3import "nx_syscalls.nx" 4import "nx_tier.nx" 5import "nx_f32.nx" 6import "nx_f32_div.nx" 7import "nx_f32_cvt.nx" 8import "nx_f32_rmsnorm.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 // Verdict gate 17 var vi: nx_int = 0 18 while vi < NX_F32_RMSN_N_VERDICTS { 19 if nx_f32_rmsnorm_verdict_is_valid(vi) != 1 { return 5 + vi } 20 vi = vi + 1 21 } 22 23 // ===== Test 1: uniform x=[1,1], gamma=[1,1] -> out=[1,1] (exact) ===== 24 let x1: *i64 = sys_mmap(2 * 8) as *i64 25 let g1: *i64 = sys_mmap(2 * 8) as *i64 26 let o1: *i64 = sys_mmap(2 * 8) as *i64 27 x1[0] = 0x3F800000; x1[1] = 0x3F800000 28 g1[0] = 0x3F800000; g1[1] = 0x3F800000 29 30 let v1: nx_int = nx_f32_rmsnorm(x1, g1, 2, 0, o1) 31 if v1 != NX_F32_RMSN_OK { return 10 } 32 // ss = 1+1 = 2, mean = 1, sqrt(1) = 1, inv = 1 33 // out[i] = 1 * 1 * 1 = 1.0 34 if o1[0] != 0x3F800000 { return 11 } 35 if o1[1] != 0x3F800000 { return 12 } 36 37 // ===== Test 2: x=[2,0], gamma=[1,1] ===== 38 // ss = 4, mean = 2, sqrt(2) ~= 1.4142, inv ~= 0.7071 39 // out[0] = 2 * 0.7071 ~= 1.4142 = 0x3FB504F3 40 // out[1] = 0 41 let x2: *i64 = sys_mmap(2 * 8) as *i64 42 let g2: *i64 = sys_mmap(2 * 8) as *i64 43 let o2: *i64 = sys_mmap(2 * 8) as *i64 44 x2[0] = 0x40000000; x2[1] = 0x00000000 45 g2[0] = 0x3F800000; g2[1] = 0x3F800000 46 47 let v2: nx_int = nx_f32_rmsnorm(x2, g2, 2, 0, o2) 48 if v2 != NX_F32_RMSN_OK { return 20 } 49 if _ulp_diff_pos(o2[0], 0x3FB504F3) > 4096 { return 21 } 50 if o2[1] != 0 { return 22 } 51 52 // ===== Test 3: x=[3,4], gamma=[1,1] ===== 53 // ss = 9 + 16 = 25, mean = 12.5, sqrt(12.5) ~= 3.535534, inv ~= 0.282843 54 // out[0] = 3 * 0.282843 = 0.848528 = 0x3F593E25 approx 55 // out[1] = 4 * 0.282843 = 1.131371 = 0x3F90C8B0 approx 56 let x3: *i64 = sys_mmap(2 * 8) as *i64 57 let g3: *i64 = sys_mmap(2 * 8) as *i64 58 let o3: *i64 = sys_mmap(2 * 8) as *i64 59 x3[0] = 0x40400000; x3[1] = 0x40800000 // 3, 4 60 g3[0] = 0x3F800000; g3[1] = 0x3F800000 61 62 let v3: nx_int = nx_f32_rmsnorm(x3, g3, 2, 0, o3) 63 if v3 != NX_F32_RMSN_OK { return 30 } 64 // 0.848528 in f32 = 0x3F593E25 -- let me double-check. Real 65 // 0.848528 has significand 0xB27C4A * 0.5 = 0.84852814... 66 // Bit pattern: exp = 126, sign = 0, mant = ~0x593E25 67 // Accept large tolerance. 68 let approx_848: i64 = 0x3F593E25 69 let approx_1131: i64 = 0x3F90D7AA 70 if _ulp_diff_pos(o3[0], approx_848) > 16384 { return 31 } 71 if _ulp_diff_pos(o3[1], approx_1131) > 16384 { return 32 } 72 73 // ===== Test 4: gamma applied as element-wise scale ===== 74 // x=[1,1], gamma=[2,3] -> out=[2, 3] (since normalizer = 1) 75 let x4: *i64 = sys_mmap(2 * 8) as *i64 76 let g4: *i64 = sys_mmap(2 * 8) as *i64 77 let o4: *i64 = sys_mmap(2 * 8) as *i64 78 x4[0] = 0x3F800000; x4[1] = 0x3F800000 // 1, 1 79 g4[0] = 0x40000000; g4[1] = 0x40400000 // 2, 3 80 81 let v4: nx_int = nx_f32_rmsnorm(x4, g4, 2, 0, o4) 82 if v4 != NX_F32_RMSN_OK { return 40 } 83 if o4[0] != 0x40000000 { return 41 } // 2.0 exact 84 if o4[1] != 0x40400000 { return 42 } // 3.0 exact 85 86 return 0 87}