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}