code wiki / (root) / nx_f32_matmul_test.nx

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}