nx_f32_matmul_test.nx source
↩ module page · 102 lines · 3887 B
1// nx_f32_matmul_test.nx -- smoke for nx_f32_matmul.nx.
2
3import "nx_syscalls.nx"
4import "nx_tier.nx"
5import "nx_f32.nx"
6import "nx_f32_matmul.nx"
7
8func _ulp_diff_pos(a: i64, b: i64) -> i64 {
9 if a >= b { return a - b }
10 return b - a
11}
12
13func main() -> i64 {
14 // Verdict gate
15 var vi: nx_int = 0
16 while vi < NX_F32_MM_N_VERDICTS {
17 if nx_f32_mm_verdict_is_valid(vi) != 1 { return 5 + vi }
18 vi = vi + 1
19 }
20
21 // ===== dot: [2.0] . [3.0] = 6.0 (exact) =====
22 let a1: *i64 = sys_mmap(8) as *i64
23 let b1: *i64 = sys_mmap(8) as *i64
24 a1[0] = 0x40000000 // 2.0
25 b1[0] = 0x40400000 // 3.0
26 let d1: i64 = nx_f32_dot(a1, b1, 1)
27 if d1 != 0x40C00000 { return 10 } // 6.0
28
29 // ===== dot: [1.0, 2.0] . [3.0, 4.0] = 3 + 8 = 11.0 =====
30 let a2: *i64 = sys_mmap(2 * 8) as *i64
31 let b2: *i64 = sys_mmap(2 * 8) as *i64
32 a2[0] = 0x3F800000; a2[1] = 0x40000000 // 1, 2
33 b2[0] = 0x40400000; b2[1] = 0x40800000 // 3, 4
34 let d2: i64 = nx_f32_dot(a2, b2, 2)
35 if d2 != 0x41300000 { return 20 } // 11.0
36
37 // ===== dot: [1, 1, 1, 1] . [1, 2, 3, 4] = 10.0 =====
38 let a3: *i64 = sys_mmap(4 * 8) as *i64
39 let b3: *i64 = sys_mmap(4 * 8) as *i64
40 a3[0] = 0x3F800000; a3[1] = 0x3F800000
41 a3[2] = 0x3F800000; a3[3] = 0x3F800000
42 b3[0] = 0x3F800000; b3[1] = 0x40000000
43 b3[2] = 0x40400000; b3[3] = 0x40800000
44 let d3: i64 = nx_f32_dot(a3, b3, 4)
45 if d3 != 0x41200000 { return 30 } // 10.0
46
47 // ===== matmul 1x2 @ 2x2 = 1x2 =====
48 // A = [[1, 2]], B = [[3, 4], [5, 6]]
49 // C = [[1*3 + 2*5, 1*4 + 2*6]] = [[13, 16]]
50 let A4: *i64 = sys_mmap(2 * 8) as *i64
51 let B4: *i64 = sys_mmap(4 * 8) as *i64
52 let C4: *i64 = sys_mmap(2 * 8) as *i64
53 A4[0] = 0x3F800000; A4[1] = 0x40000000 // 1, 2
54 B4[0] = 0x40400000; B4[1] = 0x40800000 // 3, 4
55 B4[2] = 0x40A00000; B4[3] = 0x40C00000 // 5, 6
56 let v4: nx_int = nx_f32_matmul(A4, B4, C4, 1, 2, 2)
57 if v4 != NX_F32_MM_OK { return 40 }
58 if C4[0] != 0x41500000 { return 41 } // 13.0
59 if C4[1] != 0x41800000 { return 42 } // 16.0
60
61 // ===== matmul 2x2 @ 2x2 (identity-times-something) =====
62 // A = [[1, 0], [0, 1]] (identity)
63 // B = [[5, 6], [7, 8]]
64 // C = B = [[5, 6], [7, 8]]
65 let A5: *i64 = sys_mmap(4 * 8) as *i64
66 let B5: *i64 = sys_mmap(4 * 8) as *i64
67 let C5: *i64 = sys_mmap(4 * 8) as *i64
68 A5[0] = 0x3F800000; A5[1] = 0
69 A5[2] = 0; A5[3] = 0x3F800000
70 B5[0] = 0x40A00000; B5[1] = 0x40C00000 // 5, 6
71 B5[2] = 0x40E00000; B5[3] = 0x41000000 // 7, 8
72 let v5: nx_int = nx_f32_matmul(A5, B5, C5, 2, 2, 2)
73 if v5 != NX_F32_MM_OK { return 50 }
74 if C5[0] != B5[0] { return 51 }
75 if C5[1] != B5[1] { return 52 }
76 if C5[2] != B5[2] { return 53 }
77 if C5[3] != B5[3] { return 54 }
78
79 // ===== matmul 2x3 @ 3x2 = 2x2 (general check) =====
80 // A = [[1, 2, 3], [4, 5, 6]]
81 // B = [[7, 8], [9, 10], [11, 12]]
82 // C[0,0] = 1*7 + 2*9 + 3*11 = 7+18+33 = 58
83 // C[0,1] = 1*8 + 2*10 + 3*12 = 8+20+36 = 64
84 // C[1,0] = 4*7 + 5*9 + 6*11 = 28+45+66 = 139
85 // C[1,1] = 4*8 + 5*10 + 6*12 = 32+50+72 = 154
86 let A6: *i64 = sys_mmap(6 * 8) as *i64
87 let B6: *i64 = sys_mmap(6 * 8) as *i64
88 let C6: *i64 = sys_mmap(4 * 8) as *i64
89 A6[0] = 0x3F800000; A6[1] = 0x40000000; A6[2] = 0x40400000 // 1,2,3
90 A6[3] = 0x40800000; A6[4] = 0x40A00000; A6[5] = 0x40C00000 // 4,5,6
91 B6[0] = 0x40E00000; B6[1] = 0x41000000 // 7,8
92 B6[2] = 0x41100000; B6[3] = 0x41200000 // 9,10
93 B6[4] = 0x41300000; B6[5] = 0x41400000 // 11,12
94 let v6: nx_int = nx_f32_matmul(A6, B6, C6, 2, 3, 2)
95 if v6 != NX_F32_MM_OK { return 60 }
96 if C6[0] != 0x42680000 { return 61 } // 58.0
97 if C6[1] != 0x42800000 { return 62 } // 64.0
98 if C6[2] != 0x430B0000 { return 63 } // 139.0
99 if C6[3] != 0x431A0000 { return 64 } // 154.0
100
101 return 0
102}